Support Anthropic-routed Claude models
This commit is contained in:
@@ -25,6 +25,13 @@ PROVIDER_MIN_TEMPERATURES = {
|
||||
'wenxin': 0.1,
|
||||
}
|
||||
|
||||
ANTHROPIC_MESSAGES_MODEL_PREFIXES = (
|
||||
'ppio/pa/claude-',
|
||||
)
|
||||
|
||||
ANTHROPIC_VERSION = '2023-06-01'
|
||||
ANTHROPIC_MAX_TOKENS = 4096
|
||||
|
||||
|
||||
def _join_url(base_url: str, suffix: str) -> str:
|
||||
base = base_url.rstrip('/')
|
||||
@@ -104,6 +111,24 @@ def _temperature_for_model(model: str, configured: float) -> float:
|
||||
return max(configured, minimum)
|
||||
|
||||
|
||||
def _uses_anthropic_messages_api(model: str, base_url: str) -> bool:
|
||||
normalized_model = model.strip().lower()
|
||||
normalized_base = base_url.rstrip('/').lower()
|
||||
return normalized_base.endswith('/anthropic') or any(
|
||||
normalized_model.startswith(prefix)
|
||||
for prefix in ANTHROPIC_MESSAGES_MODEL_PREFIXES
|
||||
)
|
||||
|
||||
|
||||
def _anthropic_base_url(base_url: str) -> str:
|
||||
base = base_url.rstrip('/')
|
||||
if base.lower().endswith('/anthropic'):
|
||||
return base
|
||||
if base.lower().endswith('/v1'):
|
||||
return f'{base[:-3]}/anthropic'
|
||||
return f'{base}/anthropic'
|
||||
|
||||
|
||||
def _parse_usage(payload: Any) -> UsageStats:
|
||||
if not isinstance(payload, dict):
|
||||
return UsageStats()
|
||||
@@ -160,6 +185,12 @@ class OpenAICompatClient:
|
||||
*,
|
||||
output_schema: OutputSchemaConfig | None = None,
|
||||
) -> AssistantTurn:
|
||||
if self._uses_anthropic_messages_api():
|
||||
return self._complete_anthropic_messages(
|
||||
messages=messages,
|
||||
tools=tools,
|
||||
output_schema=output_schema,
|
||||
)
|
||||
payload = self._request_json(
|
||||
self._build_payload(
|
||||
messages=messages,
|
||||
@@ -201,6 +232,13 @@ class OpenAICompatClient:
|
||||
*,
|
||||
output_schema: OutputSchemaConfig | None = None,
|
||||
) -> Iterator[StreamEvent]:
|
||||
if self._uses_anthropic_messages_api():
|
||||
yield from self._stream_anthropic_messages(
|
||||
messages=messages,
|
||||
tools=tools,
|
||||
output_schema=output_schema,
|
||||
)
|
||||
return
|
||||
payload = self._build_payload(
|
||||
messages=messages,
|
||||
tools=tools,
|
||||
@@ -231,6 +269,9 @@ class OpenAICompatClient:
|
||||
f'Unable to reach local model backend at {self.config.base_url}: {exc.reason}'
|
||||
) from exc
|
||||
|
||||
def _uses_anthropic_messages_api(self) -> bool:
|
||||
return _uses_anthropic_messages_api(self.config.model, self.config.base_url)
|
||||
|
||||
def _request_json(self, payload: dict[str, Any]) -> dict[str, Any]:
|
||||
body = json.dumps(payload).encode('utf-8')
|
||||
req = request.Request(
|
||||
@@ -287,6 +328,363 @@ class OpenAICompatClient:
|
||||
payload['response_format'] = response_format
|
||||
return payload
|
||||
|
||||
def _build_anthropic_payload(
|
||||
self,
|
||||
*,
|
||||
messages: list[dict[str, Any]],
|
||||
tools: list[dict[str, Any]],
|
||||
stream: bool,
|
||||
output_schema: OutputSchemaConfig | None,
|
||||
) -> dict[str, Any]:
|
||||
system_parts: list[str] = []
|
||||
anthropic_messages: list[dict[str, Any]] = []
|
||||
for message in messages:
|
||||
role = message.get('role')
|
||||
if role == 'system':
|
||||
content = _normalize_content(message.get('content')).strip()
|
||||
if content:
|
||||
system_parts.append(content)
|
||||
continue
|
||||
if role == 'assistant':
|
||||
blocks = self._anthropic_assistant_blocks(message)
|
||||
if blocks:
|
||||
self._append_anthropic_message(anthropic_messages, 'assistant', blocks)
|
||||
continue
|
||||
if role == 'tool':
|
||||
tool_call_id = message.get('tool_call_id')
|
||||
if not isinstance(tool_call_id, str) or not tool_call_id:
|
||||
tool_call_id = 'toolu_unknown'
|
||||
self._append_anthropic_message(
|
||||
anthropic_messages,
|
||||
'user',
|
||||
[
|
||||
{
|
||||
'type': 'tool_result',
|
||||
'tool_use_id': tool_call_id,
|
||||
'content': _normalize_content(message.get('content')),
|
||||
}
|
||||
],
|
||||
)
|
||||
continue
|
||||
if role == 'user':
|
||||
blocks = self._anthropic_text_blocks(message.get('content'))
|
||||
if blocks:
|
||||
self._append_anthropic_message(anthropic_messages, 'user', blocks)
|
||||
|
||||
payload: dict[str, Any] = {
|
||||
'model': self.config.model,
|
||||
'messages': anthropic_messages,
|
||||
'max_tokens': ANTHROPIC_MAX_TOKENS,
|
||||
'stream': stream,
|
||||
}
|
||||
if system_parts:
|
||||
payload['system'] = '\n\n'.join(system_parts)
|
||||
converted_tools = self._anthropic_tools(tools)
|
||||
if converted_tools:
|
||||
payload['tools'] = converted_tools
|
||||
if output_schema is not None:
|
||||
schema_hint = (
|
||||
f'请严格输出符合 JSON Schema `{output_schema.name}` 的 JSON,'
|
||||
'不要输出额外解释。'
|
||||
)
|
||||
payload['system'] = (
|
||||
f'{payload.get("system", "")}\n\n{schema_hint}'.strip()
|
||||
)
|
||||
return payload
|
||||
|
||||
def _anthropic_headers(self) -> dict[str, str]:
|
||||
return {
|
||||
'Authorization': f'Bearer {self.config.api_key}',
|
||||
'Content-Type': 'application/json',
|
||||
'anthropic-version': ANTHROPIC_VERSION,
|
||||
'api-key': self.config.api_key,
|
||||
'x-api-key': self.config.api_key,
|
||||
}
|
||||
|
||||
def _anthropic_messages_url(self) -> str:
|
||||
return _join_url(_anthropic_base_url(self.config.base_url), '/v1/messages')
|
||||
|
||||
def _complete_anthropic_messages(
|
||||
self,
|
||||
*,
|
||||
messages: list[dict[str, Any]],
|
||||
tools: list[dict[str, Any]],
|
||||
output_schema: OutputSchemaConfig | None,
|
||||
) -> AssistantTurn:
|
||||
payload = self._build_anthropic_payload(
|
||||
messages=messages,
|
||||
tools=tools,
|
||||
stream=False,
|
||||
output_schema=output_schema,
|
||||
)
|
||||
body = json.dumps(payload).encode('utf-8')
|
||||
req = request.Request(
|
||||
self._anthropic_messages_url(),
|
||||
data=body,
|
||||
headers=self._anthropic_headers(),
|
||||
method='POST',
|
||||
)
|
||||
try:
|
||||
with request.urlopen(req, timeout=self.config.timeout_seconds) as response:
|
||||
raw = response.read()
|
||||
except error.HTTPError as exc:
|
||||
detail = exc.read().decode('utf-8', errors='replace')
|
||||
raise OpenAICompatError(
|
||||
f'HTTP {exc.code} from local model backend: {detail}'
|
||||
) from exc
|
||||
except error.URLError as exc:
|
||||
raise OpenAICompatError(
|
||||
f'Unable to reach local model backend at {self.config.base_url}: {exc.reason}'
|
||||
) from exc
|
||||
|
||||
try:
|
||||
response_payload = json.loads(raw.decode('utf-8'))
|
||||
except json.JSONDecodeError as exc:
|
||||
raise OpenAICompatError('Local model backend returned invalid JSON') from exc
|
||||
if not isinstance(response_payload, dict):
|
||||
raise OpenAICompatError('Local model backend returned malformed JSON payload')
|
||||
return self._parse_anthropic_message_response(response_payload)
|
||||
|
||||
def _stream_anthropic_messages(
|
||||
self,
|
||||
*,
|
||||
messages: list[dict[str, Any]],
|
||||
tools: list[dict[str, Any]],
|
||||
output_schema: OutputSchemaConfig | None,
|
||||
) -> Iterator[StreamEvent]:
|
||||
payload = self._build_anthropic_payload(
|
||||
messages=messages,
|
||||
tools=tools,
|
||||
stream=True,
|
||||
output_schema=output_schema,
|
||||
)
|
||||
req = request.Request(
|
||||
self._anthropic_messages_url(),
|
||||
data=json.dumps(payload).encode('utf-8'),
|
||||
headers=self._anthropic_headers(),
|
||||
method='POST',
|
||||
)
|
||||
try:
|
||||
with request.urlopen(req, timeout=self.config.timeout_seconds) as response:
|
||||
yield StreamEvent(type='message_start')
|
||||
tool_block_indexes: dict[int, int] = {}
|
||||
next_tool_index = 0
|
||||
finish_reason: str | None = None
|
||||
for event_payload in self._iter_sse_payloads(response):
|
||||
event_type = event_payload.get('type')
|
||||
if event_type == 'message_start':
|
||||
message = event_payload.get('message')
|
||||
usage = _parse_usage(
|
||||
message.get('usage') if isinstance(message, dict) else None
|
||||
)
|
||||
if usage.total_tokens:
|
||||
yield StreamEvent(
|
||||
type='usage',
|
||||
usage=usage,
|
||||
raw_event=event_payload,
|
||||
)
|
||||
continue
|
||||
if event_type == 'content_block_start':
|
||||
block_index = event_payload.get('index')
|
||||
content_block = event_payload.get('content_block')
|
||||
if not isinstance(block_index, int) or not isinstance(content_block, dict):
|
||||
continue
|
||||
if content_block.get('type') != 'tool_use':
|
||||
continue
|
||||
tool_index = next_tool_index
|
||||
next_tool_index += 1
|
||||
tool_block_indexes[block_index] = tool_index
|
||||
tool_input = content_block.get('input')
|
||||
arguments = (
|
||||
json.dumps(tool_input, ensure_ascii=True)
|
||||
if isinstance(tool_input, dict) and tool_input
|
||||
else ''
|
||||
)
|
||||
yield StreamEvent(
|
||||
type='tool_call_delta',
|
||||
tool_call_index=tool_index,
|
||||
tool_call_id=(
|
||||
content_block.get('id')
|
||||
if isinstance(content_block.get('id'), str)
|
||||
else None
|
||||
),
|
||||
tool_name=(
|
||||
content_block.get('name')
|
||||
if isinstance(content_block.get('name'), str)
|
||||
else None
|
||||
),
|
||||
arguments_delta=arguments,
|
||||
raw_event=event_payload,
|
||||
)
|
||||
continue
|
||||
if event_type == 'content_block_delta':
|
||||
delta = event_payload.get('delta')
|
||||
if not isinstance(delta, dict):
|
||||
continue
|
||||
if delta.get('type') == 'text_delta':
|
||||
text = delta.get('text')
|
||||
if isinstance(text, str) and text:
|
||||
yield StreamEvent(
|
||||
type='content_delta',
|
||||
delta=text,
|
||||
raw_event=event_payload,
|
||||
)
|
||||
continue
|
||||
if delta.get('type') == 'input_json_delta':
|
||||
block_index = event_payload.get('index')
|
||||
partial_json = delta.get('partial_json')
|
||||
if (
|
||||
isinstance(block_index, int)
|
||||
and block_index in tool_block_indexes
|
||||
and isinstance(partial_json, str)
|
||||
and partial_json
|
||||
):
|
||||
yield StreamEvent(
|
||||
type='tool_call_delta',
|
||||
tool_call_index=tool_block_indexes[block_index],
|
||||
arguments_delta=partial_json,
|
||||
raw_event=event_payload,
|
||||
)
|
||||
continue
|
||||
if event_type == 'message_delta':
|
||||
delta = event_payload.get('delta')
|
||||
if isinstance(delta, dict) and isinstance(delta.get('stop_reason'), str):
|
||||
finish_reason = delta['stop_reason']
|
||||
usage = _parse_usage(event_payload.get('usage'))
|
||||
if usage.total_tokens:
|
||||
yield StreamEvent(
|
||||
type='usage',
|
||||
usage=usage,
|
||||
raw_event=event_payload,
|
||||
)
|
||||
continue
|
||||
if event_type == 'message_stop':
|
||||
yield StreamEvent(
|
||||
type='message_stop',
|
||||
finish_reason=finish_reason,
|
||||
raw_event=event_payload,
|
||||
)
|
||||
except error.HTTPError as exc:
|
||||
detail = exc.read().decode('utf-8', errors='replace')
|
||||
raise OpenAICompatError(
|
||||
f'HTTP {exc.code} from local model backend: {detail}'
|
||||
) from exc
|
||||
except error.URLError as exc:
|
||||
raise OpenAICompatError(
|
||||
f'Unable to reach local model backend at {self.config.base_url}: {exc.reason}'
|
||||
) from exc
|
||||
|
||||
def _parse_anthropic_message_response(self, payload: dict[str, Any]) -> AssistantTurn:
|
||||
raw_content = payload.get('content')
|
||||
if not isinstance(raw_content, list):
|
||||
raise OpenAICompatError('Anthropic backend returned no content blocks')
|
||||
text_parts: list[str] = []
|
||||
tool_calls: list[ToolCall] = []
|
||||
for index, block in enumerate(raw_content):
|
||||
if not isinstance(block, dict):
|
||||
continue
|
||||
if block.get('type') == 'text':
|
||||
text = block.get('text')
|
||||
if isinstance(text, str):
|
||||
text_parts.append(text)
|
||||
continue
|
||||
if block.get('type') == 'tool_use':
|
||||
name = block.get('name')
|
||||
if not isinstance(name, str) or not name:
|
||||
raise OpenAICompatError('Tool call missing function name')
|
||||
call_id = block.get('id')
|
||||
if not isinstance(call_id, str) or not call_id:
|
||||
call_id = f'call_{index}'
|
||||
arguments = block.get('input')
|
||||
if not isinstance(arguments, dict):
|
||||
arguments = {}
|
||||
tool_calls.append(ToolCall(id=call_id, name=name, arguments=arguments))
|
||||
finish_reason = payload.get('stop_reason')
|
||||
if finish_reason is not None and not isinstance(finish_reason, str):
|
||||
finish_reason = str(finish_reason)
|
||||
return AssistantTurn(
|
||||
content=''.join(text_parts),
|
||||
tool_calls=tuple(tool_calls),
|
||||
finish_reason=finish_reason,
|
||||
raw_message=payload,
|
||||
usage=_parse_usage(payload.get('usage')),
|
||||
)
|
||||
|
||||
def _anthropic_assistant_blocks(self, message: dict[str, Any]) -> list[dict[str, Any]]:
|
||||
blocks: list[dict[str, Any]] = []
|
||||
content = _normalize_content(message.get('content'))
|
||||
if content:
|
||||
blocks.append({'type': 'text', 'text': content})
|
||||
raw_tool_calls = message.get('tool_calls')
|
||||
if isinstance(raw_tool_calls, list):
|
||||
for index, raw_call in enumerate(raw_tool_calls):
|
||||
if not isinstance(raw_call, dict):
|
||||
continue
|
||||
function_block = raw_call.get('function')
|
||||
if not isinstance(function_block, dict):
|
||||
continue
|
||||
name = function_block.get('name')
|
||||
if not isinstance(name, str) or not name:
|
||||
continue
|
||||
call_id = raw_call.get('id')
|
||||
if not isinstance(call_id, str) or not call_id:
|
||||
call_id = f'call_{index}'
|
||||
try:
|
||||
arguments = _parse_tool_arguments(function_block.get('arguments'))
|
||||
except OpenAICompatError:
|
||||
arguments = {}
|
||||
blocks.append(
|
||||
{
|
||||
'type': 'tool_use',
|
||||
'id': call_id,
|
||||
'name': name,
|
||||
'input': arguments,
|
||||
}
|
||||
)
|
||||
return blocks
|
||||
|
||||
def _anthropic_text_blocks(self, content: Any) -> list[dict[str, Any]]:
|
||||
normalized = _normalize_content(content)
|
||||
if not normalized:
|
||||
return []
|
||||
return [{'type': 'text', 'text': normalized}]
|
||||
|
||||
def _anthropic_tools(self, tools: list[dict[str, Any]]) -> list[dict[str, Any]]:
|
||||
converted: list[dict[str, Any]] = []
|
||||
for tool in tools:
|
||||
if not isinstance(tool, dict):
|
||||
continue
|
||||
function_block = tool.get('function')
|
||||
if not isinstance(function_block, dict):
|
||||
continue
|
||||
name = function_block.get('name')
|
||||
if not isinstance(name, str) or not name:
|
||||
continue
|
||||
entry: dict[str, Any] = {
|
||||
'name': name,
|
||||
'input_schema': function_block.get('parameters') or {'type': 'object'},
|
||||
}
|
||||
description = function_block.get('description')
|
||||
if isinstance(description, str) and description:
|
||||
entry['description'] = description
|
||||
converted.append(entry)
|
||||
return converted
|
||||
|
||||
def _append_anthropic_message(
|
||||
self,
|
||||
messages: list[dict[str, Any]],
|
||||
role: str,
|
||||
blocks: list[dict[str, Any]],
|
||||
) -> None:
|
||||
if not blocks:
|
||||
return
|
||||
if messages and messages[-1].get('role') == role:
|
||||
previous = messages[-1].get('content')
|
||||
if isinstance(previous, list):
|
||||
previous.extend(blocks)
|
||||
return
|
||||
messages.append({'role': role, 'content': blocks})
|
||||
|
||||
def _parse_tool_calls_from_message(self, message: dict[str, Any]) -> list[ToolCall]:
|
||||
tool_calls: list[ToolCall] = []
|
||||
raw_tool_calls = message.get('tool_calls')
|
||||
|
||||
Reference in New Issue
Block a user