435 lines
16 KiB
Python
435 lines
16 KiB
Python
from __future__ import annotations
|
|
|
|
import json
|
|
import unittest
|
|
from unittest.mock import patch
|
|
|
|
from src.openai_compat import (
|
|
OpenAICompatClient,
|
|
OpenAICompatError,
|
|
_anthropic_base_url,
|
|
_build_response_format,
|
|
_join_url,
|
|
_normalize_content,
|
|
_optional_int,
|
|
_parse_tool_arguments,
|
|
_parse_usage,
|
|
_temperature_for_model,
|
|
_uses_anthropic_messages_api,
|
|
)
|
|
from src.agent_types import ModelConfig, OutputSchemaConfig, UsageStats
|
|
|
|
|
|
class FakeHTTPResponse:
|
|
def __init__(self, payload: dict[str, object]) -> None:
|
|
self.payload = payload
|
|
|
|
def read(self) -> bytes:
|
|
return json.dumps(self.payload).encode('utf-8')
|
|
|
|
def __enter__(self) -> 'FakeHTTPResponse':
|
|
return self
|
|
|
|
def __exit__(self, exc_type, exc, tb) -> None:
|
|
return None
|
|
|
|
|
|
class FakeStreamingHTTPResponse:
|
|
def __init__(self, payloads: list[dict[str, object]]) -> None:
|
|
self.lines: list[bytes] = []
|
|
for payload in payloads:
|
|
self.lines.append(b'event: message\n')
|
|
self.lines.append(f'data: {json.dumps(payload)}\n'.encode('utf-8'))
|
|
self.lines.append(b'\n')
|
|
|
|
def readline(self) -> bytes:
|
|
if not self.lines:
|
|
return b''
|
|
return self.lines.pop(0)
|
|
|
|
def __enter__(self) -> 'FakeStreamingHTTPResponse':
|
|
return self
|
|
|
|
def __exit__(self, exc_type, exc, tb) -> None:
|
|
return None
|
|
|
|
|
|
class TestJoinUrl(unittest.TestCase):
|
|
def test_base_with_trailing_slash(self):
|
|
self.assertEqual(_join_url('http://localhost:8000/', 'v1/chat'), 'http://localhost:8000/v1/chat')
|
|
|
|
def test_base_without_trailing_slash(self):
|
|
self.assertEqual(_join_url('http://localhost:8000', 'v1/chat'), 'http://localhost:8000/v1/chat')
|
|
|
|
def test_suffix_with_leading_slash(self):
|
|
self.assertEqual(_join_url('http://localhost:8000', '/v1/chat'), 'http://localhost:8000/v1/chat')
|
|
|
|
|
|
class TestNormalizeContent(unittest.TestCase):
|
|
def test_string_passthrough(self):
|
|
self.assertEqual(_normalize_content('hello'), 'hello')
|
|
|
|
def test_none_returns_empty(self):
|
|
self.assertEqual(_normalize_content(None), '')
|
|
|
|
def test_list_of_strings_joined(self):
|
|
self.assertEqual(_normalize_content(['hello', ' ', 'world']), 'hello world')
|
|
|
|
def test_list_of_text_dicts(self):
|
|
items = [{'type': 'text', 'text': 'hello'}, {'type': 'text', 'text': ' world'}]
|
|
self.assertEqual(_normalize_content(items), 'hello world')
|
|
|
|
def test_list_of_mixed_items(self):
|
|
items = ['start ', {'type': 'text', 'text': 'middle'}, ' end']
|
|
self.assertEqual(_normalize_content(items), 'start middle end')
|
|
|
|
def test_non_string_non_list_returns_str(self):
|
|
self.assertEqual(_normalize_content(42), '42')
|
|
|
|
|
|
class TestParseToolArguments(unittest.TestCase):
|
|
def test_dict_passthrough(self):
|
|
d = {'key': 'value'}
|
|
self.assertIs(_parse_tool_arguments(d), d)
|
|
|
|
def test_valid_json_string(self):
|
|
self.assertEqual(_parse_tool_arguments('{"a": 1}'), {'a': 1})
|
|
|
|
def test_empty_string_returns_empty_dict(self):
|
|
self.assertEqual(_parse_tool_arguments(''), {})
|
|
|
|
def test_none_returns_empty_dict(self):
|
|
self.assertEqual(_parse_tool_arguments(None), {})
|
|
|
|
def test_invalid_json_raises(self):
|
|
with self.assertRaises(OpenAICompatError):
|
|
_parse_tool_arguments('{bad json}')
|
|
|
|
def test_json_non_dict_raises(self):
|
|
with self.assertRaises(OpenAICompatError):
|
|
_parse_tool_arguments('[1, 2, 3]')
|
|
|
|
def test_unsupported_type_raises(self):
|
|
with self.assertRaises(OpenAICompatError):
|
|
_parse_tool_arguments(12345)
|
|
|
|
|
|
class TestParseUsage(unittest.TestCase):
|
|
def test_standard_fields(self):
|
|
usage = _parse_usage({'input_tokens': 10, 'output_tokens': 20})
|
|
self.assertEqual(usage.input_tokens, 10)
|
|
self.assertEqual(usage.output_tokens, 20)
|
|
|
|
def test_prompt_completion_aliases(self):
|
|
usage = _parse_usage({'prompt_tokens': 15, 'completion_tokens': 25})
|
|
self.assertEqual(usage.input_tokens, 15)
|
|
self.assertEqual(usage.output_tokens, 25)
|
|
|
|
def test_ollama_aliases(self):
|
|
usage = _parse_usage({'prompt_eval_count': 12, 'eval_count': 18})
|
|
self.assertEqual(usage.input_tokens, 12)
|
|
self.assertEqual(usage.output_tokens, 18)
|
|
|
|
def test_cache_tokens(self):
|
|
usage = _parse_usage({
|
|
'input_tokens': 1,
|
|
'output_tokens': 1,
|
|
'cache_creation_input_tokens': 100,
|
|
'cache_read_input_tokens': 200,
|
|
})
|
|
self.assertEqual(usage.cache_creation_input_tokens, 100)
|
|
self.assertEqual(usage.cache_read_input_tokens, 200)
|
|
|
|
def test_reasoning_tokens_top_level_and_details(self):
|
|
usage_top = _parse_usage({'input_tokens': 1, 'output_tokens': 1, 'reasoning_tokens': 50})
|
|
self.assertEqual(usage_top.reasoning_tokens, 50)
|
|
|
|
usage_details = _parse_usage({
|
|
'input_tokens': 1,
|
|
'output_tokens': 1,
|
|
'completion_tokens_details': {'reasoning_tokens': 75},
|
|
})
|
|
self.assertEqual(usage_details.reasoning_tokens, 75)
|
|
|
|
def test_non_dict_returns_empty(self):
|
|
usage = _parse_usage('not a dict')
|
|
self.assertEqual(usage, UsageStats())
|
|
|
|
def test_string_number_coercion(self):
|
|
usage = _parse_usage({'input_tokens': '10', 'output_tokens': '20'})
|
|
self.assertEqual(usage.input_tokens, 10)
|
|
self.assertEqual(usage.output_tokens, 20)
|
|
|
|
|
|
class TestBuildResponseFormat(unittest.TestCase):
|
|
def test_none_returns_none(self):
|
|
self.assertIsNone(_build_response_format(None))
|
|
|
|
def test_valid_schema(self):
|
|
schema = OutputSchemaConfig(
|
|
name='test_schema',
|
|
schema={'type': 'object', 'properties': {'x': {'type': 'integer'}}},
|
|
strict=True,
|
|
)
|
|
result = _build_response_format(schema)
|
|
self.assertEqual(result, {
|
|
'type': 'json_schema',
|
|
'json_schema': {
|
|
'name': 'test_schema',
|
|
'schema': {'type': 'object', 'properties': {'x': {'type': 'integer'}}},
|
|
'strict': True,
|
|
},
|
|
})
|
|
|
|
|
|
class TestOptionalInt(unittest.TestCase):
|
|
def test_int_passthrough(self):
|
|
self.assertEqual(_optional_int(42), 42)
|
|
|
|
def test_float_truncated(self):
|
|
self.assertEqual(_optional_int(3.9), 3)
|
|
|
|
def test_string_parsed(self):
|
|
self.assertEqual(_optional_int('7'), 7)
|
|
|
|
def test_bool_returns_zero(self):
|
|
self.assertEqual(_optional_int(True), 0)
|
|
self.assertEqual(_optional_int(False), 0)
|
|
|
|
def test_none_returns_zero(self):
|
|
self.assertEqual(_optional_int(None), 0)
|
|
|
|
def test_invalid_string_returns_zero(self):
|
|
self.assertEqual(_optional_int('abc'), 0)
|
|
|
|
|
|
class TestProviderTemperatureFloor(unittest.TestCase):
|
|
def test_minimax_temperature_is_clamped_to_provider_minimum(self):
|
|
self.assertEqual(_temperature_for_model('minimax/MiniMax-M2.5', 0.0), 0.01)
|
|
|
|
def test_wenxin_temperature_is_clamped_to_provider_minimum(self):
|
|
self.assertEqual(_temperature_for_model('wenxin/ernie-4.0-turbo-128k', 0.0), 0.1)
|
|
|
|
def test_regular_models_keep_configured_temperature(self):
|
|
self.assertEqual(_temperature_for_model('xiaomi/mimo-v2-flash', 0.0), 0.0)
|
|
|
|
def test_payload_uses_temperature_floor(self):
|
|
client = OpenAICompatClient(
|
|
ModelConfig(model='minimax/MiniMax-M2.5', temperature=0.0)
|
|
)
|
|
payload = client._build_payload( # noqa: SLF001 - verify payload compatibility.
|
|
messages=[{'role': 'user', 'content': 'hi'}],
|
|
tools=[],
|
|
stream=False,
|
|
output_schema=None,
|
|
)
|
|
self.assertEqual(payload['temperature'], 0.01)
|
|
|
|
def test_payload_omits_tool_choice_when_no_tools_are_supplied(self):
|
|
client = OpenAICompatClient(ModelConfig(model='azure_openai/gpt-4o-mini'))
|
|
payload = client._build_payload( # noqa: SLF001 - verify provider compatibility.
|
|
messages=[{'role': 'user', 'content': 'hi'}],
|
|
tools=[],
|
|
stream=False,
|
|
output_schema=None,
|
|
)
|
|
self.assertNotIn('tools', payload)
|
|
self.assertNotIn('tool_choice', payload)
|
|
|
|
def test_payload_includes_tool_choice_when_tools_are_supplied(self):
|
|
client = OpenAICompatClient(ModelConfig(model='azure_openai/gpt-4o-mini'))
|
|
payload = client._build_payload( # noqa: SLF001 - verify provider compatibility.
|
|
messages=[{'role': 'user', 'content': 'hi'}],
|
|
tools=[{'type': 'function', 'function': {'name': 'noop', 'parameters': {}}}],
|
|
stream=False,
|
|
output_schema=None,
|
|
)
|
|
self.assertIn('tools', payload)
|
|
self.assertEqual(payload['tool_choice'], 'auto')
|
|
|
|
|
|
class TestAnthropicMessagesRouting(unittest.TestCase):
|
|
def test_ppio_pa_claude_uses_anthropic_messages_api(self):
|
|
self.assertTrue(
|
|
_uses_anthropic_messages_api(
|
|
'ppio/pa/claude-opus-4-7',
|
|
'http://model.mify.ai.srv/v1',
|
|
)
|
|
)
|
|
self.assertFalse(
|
|
_uses_anthropic_messages_api(
|
|
'ppio/gemini-2.5-pro',
|
|
'http://model.mify.ai.srv/v1',
|
|
)
|
|
)
|
|
|
|
def test_anthropic_base_url_is_derived_from_openai_base(self):
|
|
self.assertEqual(
|
|
_anthropic_base_url('http://model.mify.ai.srv/v1'),
|
|
'http://model.mify.ai.srv/anthropic',
|
|
)
|
|
self.assertEqual(
|
|
_anthropic_base_url('http://model.mify.ai.srv/anthropic'),
|
|
'http://model.mify.ai.srv/anthropic',
|
|
)
|
|
|
|
def test_anthropic_complete_converts_tools_and_parses_tool_use(self):
|
|
recorded: dict[str, object] = {}
|
|
|
|
def fake_urlopen(request_obj, timeout=None): # noqa: ANN001
|
|
recorded['url'] = request_obj.full_url
|
|
recorded['payload'] = json.loads(request_obj.data.decode('utf-8'))
|
|
return FakeHTTPResponse(
|
|
{
|
|
'id': 'msg_1',
|
|
'type': 'message',
|
|
'role': 'assistant',
|
|
'content': [
|
|
{'type': 'text', 'text': '我来读取文件。'},
|
|
{
|
|
'type': 'tool_use',
|
|
'id': 'toolu_1',
|
|
'name': 'read_file',
|
|
'input': {'path': 'hello.txt'},
|
|
},
|
|
],
|
|
'stop_reason': 'tool_use',
|
|
'usage': {'input_tokens': 10, 'output_tokens': 4},
|
|
}
|
|
)
|
|
|
|
client = OpenAICompatClient(
|
|
ModelConfig(
|
|
model='ppio/pa/claude-opus-4-7',
|
|
base_url='http://model.mify.ai.srv/v1',
|
|
api_key='token',
|
|
)
|
|
)
|
|
with patch('src.openai_compat.request.urlopen', side_effect=fake_urlopen):
|
|
turn = client.complete(
|
|
messages=[
|
|
{'role': 'system', 'content': '你是工具型助手。'},
|
|
{'role': 'user', 'content': '读取 hello.txt'},
|
|
],
|
|
tools=[
|
|
{
|
|
'type': 'function',
|
|
'function': {
|
|
'name': 'read_file',
|
|
'description': '读取文件',
|
|
'parameters': {
|
|
'type': 'object',
|
|
'properties': {'path': {'type': 'string'}},
|
|
'required': ['path'],
|
|
},
|
|
},
|
|
}
|
|
],
|
|
)
|
|
|
|
self.assertEqual(
|
|
recorded['url'],
|
|
'http://model.mify.ai.srv/anthropic/v1/messages',
|
|
)
|
|
payload = recorded['payload']
|
|
self.assertEqual(payload['model'], 'ppio/pa/claude-opus-4-7')
|
|
self.assertEqual(payload['system'], '你是工具型助手。')
|
|
self.assertEqual(payload['tools'][0]['name'], 'read_file')
|
|
self.assertEqual(payload['tools'][0]['input_schema']['required'], ['path'])
|
|
self.assertEqual(turn.content, '我来读取文件。')
|
|
self.assertEqual(turn.finish_reason, 'tool_use')
|
|
self.assertEqual(turn.tool_calls[0].id, 'toolu_1')
|
|
self.assertEqual(turn.tool_calls[0].arguments, {'path': 'hello.txt'})
|
|
self.assertEqual(turn.usage.input_tokens, 10)
|
|
|
|
def test_anthropic_stream_parses_text_usage_and_tool_use(self):
|
|
payloads = [
|
|
{
|
|
'type': 'message_start',
|
|
'message': {'usage': {'input_tokens': 7, 'output_tokens': 1}},
|
|
},
|
|
{
|
|
'type': 'content_block_start',
|
|
'index': 0,
|
|
'content_block': {'type': 'text', 'text': ''},
|
|
},
|
|
{
|
|
'type': 'content_block_delta',
|
|
'index': 0,
|
|
'delta': {'type': 'text_delta', 'text': '读取'},
|
|
},
|
|
{
|
|
'type': 'content_block_start',
|
|
'index': 1,
|
|
'content_block': {
|
|
'type': 'tool_use',
|
|
'id': 'toolu_1',
|
|
'name': 'read_file',
|
|
'input': {},
|
|
},
|
|
},
|
|
{
|
|
'type': 'content_block_delta',
|
|
'index': 1,
|
|
'delta': {'type': 'input_json_delta', 'partial_json': '{"path":'},
|
|
},
|
|
{
|
|
'type': 'content_block_delta',
|
|
'index': 1,
|
|
'delta': {'type': 'input_json_delta', 'partial_json': '"hello.txt"}'},
|
|
},
|
|
{
|
|
'type': 'message_delta',
|
|
'delta': {'stop_reason': 'tool_use'},
|
|
'usage': {'output_tokens': 5},
|
|
},
|
|
{'type': 'message_stop'},
|
|
]
|
|
|
|
def fake_urlopen(request_obj, timeout=None): # noqa: ANN001
|
|
return FakeStreamingHTTPResponse(payloads)
|
|
|
|
client = OpenAICompatClient(
|
|
ModelConfig(
|
|
model='ppio/pa/claude-opus-4-7',
|
|
base_url='http://model.mify.ai.srv/v1',
|
|
api_key='token',
|
|
)
|
|
)
|
|
with patch('src.openai_compat.request.urlopen', side_effect=fake_urlopen):
|
|
events = list(
|
|
client.stream(
|
|
messages=[{'role': 'user', 'content': '读取 hello.txt'}],
|
|
tools=[
|
|
{
|
|
'type': 'function',
|
|
'function': {
|
|
'name': 'read_file',
|
|
'parameters': {'type': 'object'},
|
|
},
|
|
}
|
|
],
|
|
)
|
|
)
|
|
|
|
self.assertEqual(events[0].type, 'message_start')
|
|
self.assertEqual(
|
|
''.join(event.delta for event in events if event.type == 'content_delta'),
|
|
'读取',
|
|
)
|
|
tool_events = [event for event in events if event.type == 'tool_call_delta']
|
|
self.assertEqual(tool_events[0].tool_call_index, 0)
|
|
self.assertEqual(tool_events[0].tool_call_id, 'toolu_1')
|
|
self.assertEqual(tool_events[0].tool_name, 'read_file')
|
|
self.assertEqual(
|
|
''.join(event.arguments_delta for event in tool_events),
|
|
'{"path":"hello.txt"}',
|
|
)
|
|
self.assertTrue(any(event.type == 'usage' for event in events))
|
|
self.assertEqual(events[-1].type, 'message_stop')
|
|
self.assertEqual(events[-1].finish_reason, 'tool_use')
|
|
|
|
|
|
if __name__ == '__main__':
|
|
unittest.main()
|