Files
cdsl-cad/backend/tests/test_structured_llm.py
T
2026-09-04 11:17:36 +08:00

124 lines
5.2 KiB
Python

from __future__ import annotations
import asyncio
from pathlib import Path
import sys
import unittest
ROOT = Path(__file__).resolve().parents[2]
sys.path.insert(0, str(ROOT / "backend"))
from app.cad_agent.adapters.structured_llm import StructuredModelError, StructuredModelGateway # noqa: E402
from app.settings import ProviderConfig, ProviderModel, Settings # noqa: E402
def _settings(*, api_style: str = "responses", reasoning_effort: str = "medium") -> Settings:
provider = ProviderConfig(
"provider", "Provider", "https://example.invalid/v1", "key",
(ProviderModel("model"),), reasoning_effort=reasoning_effort, api_style=api_style,
)
return Settings(
task_root=ROOT / "tmp-tasks",
conversation_root=ROOT / "tmp-conversations",
library_root=ROOT / "backend" / "cdsl_library",
engine_root=ROOT / "backend" / "engine" / "cdsl_engine",
llm_base_url=provider.base_url,
llm_api_key=provider.api_key,
llm_model="model",
llm_timeout_s=1,
default_provider_id="provider",
providers=(provider,),
)
def _tool() -> dict:
return {
"type": "function",
"function": {
"name": "write_document",
"description": "Write one document.",
"parameters": {"type": "object", "properties": {}, "additionalProperties": False},
},
}
def _response(api_style: str) -> dict:
if api_style == "responses":
return {
"output": [{"type": "function_call", "name": "write_document", "arguments": "{}"}],
"usage": {"input_tokens": 5, "output_tokens": 3, "total_tokens": 8},
}
return {
"choices": [{"message": {"tool_calls": [{"function": {"name": "write_document", "arguments": "{}"}}]}}],
"usage": {"prompt_tokens": 5, "completion_tokens": 3, "total_tokens": 8},
}
class StructuredModelGatewayCompatibilityTests(unittest.TestCase):
def test_thinking_tool_choice_retries_without_reasoning_before_relaxing_tool_choice(self) -> None:
gateway = StructuredModelGateway(_settings())
payloads: list[dict] = []
async def request(_provider: object, payload: dict) -> dict:
payloads.append(payload)
if len(payloads) == 1:
raise StructuredModelError("Provider rejected structured request (400): Thinking mode does not support this tool_choice")
return _response("responses")
gateway._request = request # type: ignore[method-assign]
result = asyncio.run(gateway.call_tool(
messages=[{"role": "system", "content": "Call the tool."}],
tool=_tool(), provider_id="provider", model_id="model", required_tool_name="write_document",
))
self.assertEqual(len(payloads), 2)
self.assertEqual(payloads[0]["tool_choice"], {"type": "function", "name": "write_document"})
self.assertEqual(payloads[0]["reasoning"], {"effort": "medium"})
self.assertEqual(payloads[1]["tool_choice"], {"type": "function", "name": "write_document"})
self.assertNotIn("reasoning", payloads[1])
self.assertEqual(result["usage"]["structured_compatibility_mode"], "reasoning_disabled")
def test_persistent_thinking_rejection_uses_auto_with_the_same_single_tool(self) -> None:
gateway = StructuredModelGateway(_settings(api_style="chat_completions", reasoning_effort=""))
payloads: list[dict] = []
async def request(_provider: object, payload: dict) -> dict:
payloads.append(payload)
if len(payloads) < 3:
raise StructuredModelError("Thinking mode does not support this tool_choice")
return _response("chat_completions")
gateway._request = request # type: ignore[method-assign]
result = asyncio.run(gateway.call_tool(
messages=[{"role": "system", "content": "Call the tool."}],
tool=_tool(), provider_id="provider", model_id="model", required_tool_name="write_document",
))
self.assertEqual(len(payloads), 3)
self.assertEqual(payloads[0]["tool_choice"]["function"]["name"], "write_document")
self.assertEqual(payloads[1]["tool_choice"]["function"]["name"], "write_document")
self.assertEqual(payloads[2]["tool_choice"], "auto")
self.assertEqual(len(payloads[2]["tools"]), 1)
self.assertEqual(result["usage"]["structured_compatibility_mode"], "single_tool_auto")
def test_unrelated_provider_rejection_is_not_retried(self) -> None:
gateway = StructuredModelGateway(_settings())
payloads: list[dict] = []
async def request(_provider: object, payload: dict) -> dict:
payloads.append(payload)
raise StructuredModelError("Provider rejected structured request (400): invalid model")
gateway._request = request # type: ignore[method-assign]
with self.assertRaisesRegex(StructuredModelError, "invalid model"):
asyncio.run(gateway.call_tool(
messages=[{"role": "system", "content": "Call the tool."}],
tool=_tool(), provider_id="provider", model_id="model", required_tool_name="write_document",
))
self.assertEqual(len(payloads), 1)
if __name__ == "__main__":
unittest.main()