change
This commit is contained in:
@@ -0,0 +1,337 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import tempfile
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
from app.models.contracts import ChatMessage, MessagePart
|
||||
from app.services.agent_service import AgentService, CDSL_TOOL_SCHEMA, RepeatedToolArgumentsError, StrictToolSchemaError, TOOL_SCHEMAS, ToolArgumentsError, parse_tool_arguments, response_language_instruction, tools_for_model, user_visible_error_message
|
||||
from app.services.library import CdslLibrary
|
||||
from app.services.storage import WorkspaceStore
|
||||
from app.settings import ProviderConfig, ProviderModel, Settings
|
||||
|
||||
|
||||
class ParseToolArgumentsTests(unittest.TestCase):
|
||||
def test_accepts_one_json_object(self) -> None:
|
||||
payload = parse_tool_arguments(' {"summary":"water cup","cdsl":{"parts":[]}} ')
|
||||
|
||||
self.assertEqual(payload["summary"], "water cup")
|
||||
self.assertEqual(payload["cdsl"], {"parts": []})
|
||||
|
||||
def test_rejects_concatenated_json_objects(self) -> None:
|
||||
with self.assertRaisesRegex(ToolArgumentsError, "trailing content"):
|
||||
parse_tool_arguments('{"summary":"water cup"}{"cdsl":{}}')
|
||||
|
||||
def test_rejects_markdown_or_prose_after_json(self) -> None:
|
||||
with self.assertRaisesRegex(ToolArgumentsError, "trailing content"):
|
||||
parse_tool_arguments('{"summary":"water cup"}\n```')
|
||||
|
||||
def test_rejects_non_object_json(self) -> None:
|
||||
with self.assertRaisesRegex(ToolArgumentsError, "JSON object"):
|
||||
parse_tool_arguments('["not", "tool arguments"]')
|
||||
|
||||
def test_recovers_only_the_known_premature_cdsl_wrapper_close(self) -> None:
|
||||
payload = parse_tool_arguments(
|
||||
'{"cdsl":{"schema":"cad.cdsl.llm.v1"}}, "summary":"fixed envelope"}',
|
||||
recover_cdsl_wrapper=True,
|
||||
)
|
||||
|
||||
self.assertEqual(payload["summary"], "fixed envelope")
|
||||
self.assertEqual(payload["cdsl"], {"schema": "cad.cdsl.llm.v1"})
|
||||
|
||||
def test_does_not_recover_arbitrary_trailing_tool_content(self) -> None:
|
||||
with self.assertRaisesRegex(ToolArgumentsError, "trailing content"):
|
||||
parse_tool_arguments(
|
||||
'{"cdsl":{"schema":"cad.cdsl.llm.v1"}} prose',
|
||||
recover_cdsl_wrapper=True,
|
||||
)
|
||||
|
||||
def test_identifies_chinese_output_requirement(self) -> None:
|
||||
self.assertIn("Chinese", response_language_instruction("生成一个水杯"))
|
||||
|
||||
def test_generate_tool_requires_a_non_empty_cdsl_structure(self) -> None:
|
||||
generate_tool = next(tool for tool in TOOL_SCHEMAS if tool["function"]["name"] == "generate_cdsl_model")
|
||||
cdsl = generate_tool["function"]["parameters"]["properties"]["cdsl"]
|
||||
|
||||
self.assertEqual(set(cdsl["required"]), {"schema", "features", "geometry"})
|
||||
self.assertEqual(cdsl["properties"]["features"]["minItems"], 1)
|
||||
self.assertEqual(cdsl["properties"]["geometry"]["properties"]["sketches"]["minItems"], 1)
|
||||
self.assertIn("extrude_add_blind", cdsl["$defs"]["feature_atomic_ids"]["enum"])
|
||||
self.assertNotIn("extrude", cdsl["$defs"]["feature_atomic_ids"]["enum"])
|
||||
self.assertEqual(cdsl, CDSL_TOOL_SCHEMA)
|
||||
|
||||
def test_strict_tool_schema_is_limited_to_cdsl_generation_arguments(self) -> None:
|
||||
tools = tools_for_model(ProviderModel("strict-model", strict_tool_schema=True))
|
||||
strict_tools = [tool["function"]["name"] for tool in tools if tool["function"].get("strict")]
|
||||
|
||||
self.assertEqual(strict_tools, ["generate_cdsl_model"])
|
||||
generate_tool = next(tool for tool in tools if tool["function"]["name"] == "generate_cdsl_model")
|
||||
self.assertEqual(generate_tool["function"]["parameters"]["properties"]["summary"], {"type": "string", "minLength": 1})
|
||||
self.assertEqual(generate_tool["function"]["parameters"]["properties"]["cdsl"], CDSL_TOOL_SCHEMA)
|
||||
self.assertFalse(generate_tool["function"]["parameters"]["additionalProperties"])
|
||||
|
||||
def test_default_model_does_not_receive_strict_tool_schema(self) -> None:
|
||||
tools = tools_for_model(ProviderModel("default-model"))
|
||||
|
||||
self.assertFalse(any(tool["function"].get("strict") for tool in tools))
|
||||
|
||||
def test_strict_schema_rejection_is_localized_for_chinese_requests(self) -> None:
|
||||
message = user_visible_error_message(
|
||||
StrictToolSchemaError("provider rejected strict schema"),
|
||||
"生成一个法兰",
|
||||
)
|
||||
|
||||
self.assertIn("不支持严格 CDSL 工具 schema", message)
|
||||
self.assertNotIn("provider rejected", message)
|
||||
|
||||
def test_repeated_invalid_cdsl_tool_arguments_are_localized_for_chinese_requests(self) -> None:
|
||||
message = user_visible_error_message(
|
||||
RepeatedToolArgumentsError("arguments are not valid JSON"),
|
||||
"生成一个法兰",
|
||||
)
|
||||
|
||||
self.assertIn("连续两次未返回完整的 CDSL 工具 JSON", message)
|
||||
self.assertIn("函数调用兼容性", message)
|
||||
|
||||
|
||||
class ToolArgumentsRetryTests(unittest.TestCase):
|
||||
def test_repeated_invalid_cdsl_arguments_stop_before_the_safety_limit(self) -> None:
|
||||
class InvalidCdslAgent(AgentService):
|
||||
def __init__(self, *args: object, **kwargs: object) -> None:
|
||||
super().__init__(*args, **kwargs)
|
||||
self.responses = [
|
||||
{
|
||||
"choices": [{"message": {
|
||||
"role": "assistant",
|
||||
"content": "",
|
||||
"tool_calls": [{
|
||||
"id": "invalid_cdsl_1",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "generate_cdsl_model",
|
||||
"arguments": '{"cdsl":',
|
||||
},
|
||||
}],
|
||||
}}],
|
||||
},
|
||||
{
|
||||
"choices": [{"message": {
|
||||
"role": "assistant",
|
||||
"content": "",
|
||||
"tool_calls": [{
|
||||
"id": "invalid_cdsl_2",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "generate_cdsl_model",
|
||||
"arguments": '{"cdsl":',
|
||||
},
|
||||
}],
|
||||
}}],
|
||||
},
|
||||
]
|
||||
|
||||
self.responses[0]["id"] = "chatcmpl_invalid_1"
|
||||
self.responses[0]["model"] = "test-model"
|
||||
self.responses[0]["usage"] = {"completion_tokens": 4096}
|
||||
self.responses[0]["choices"][0]["finish_reason"] = "length"
|
||||
self.responses[1]["id"] = "chatcmpl_invalid_2"
|
||||
self.responses[1]["model"] = "test-model"
|
||||
self.responses[1]["usage"] = {"completion_tokens": 4096}
|
||||
self.responses[1]["choices"][0]["finish_reason"] = "length"
|
||||
|
||||
async def _complete(self, *args: object, **kwargs: object) -> dict[str, object]:
|
||||
return self.responses.pop(0)
|
||||
|
||||
backend_root = Path(__file__).resolve().parents[1]
|
||||
with tempfile.TemporaryDirectory() as temporary_directory:
|
||||
temporary_root = Path(temporary_directory)
|
||||
provider = ProviderConfig("test", "Test", "https://example.invalid/v1", "test-key", (ProviderModel("test-model"),))
|
||||
settings = Settings(
|
||||
task_root=temporary_root / "tasks",
|
||||
conversation_root=temporary_root / "conversations",
|
||||
library_root=backend_root / "cdsl_library",
|
||||
engine_root=backend_root / "engine" / "cdsl_engine",
|
||||
llm_base_url=provider.base_url,
|
||||
llm_api_key=provider.api_key,
|
||||
llm_model="test-model",
|
||||
llm_timeout_s=1,
|
||||
default_provider_id="test",
|
||||
providers=(provider,),
|
||||
)
|
||||
agent = InvalidCdslAgent(settings, WorkspaceStore(settings), CdslLibrary(settings))
|
||||
message = ChatMessage(id="user_1", role="user", parts=[MessagePart(type="text", text="生成一个法兰")])
|
||||
|
||||
async def collect_events() -> list[dict[str, object]]:
|
||||
events: list[dict[str, object]] = []
|
||||
async for chunk in agent.stream([message], None, None):
|
||||
events.append(json.loads(chunk.decode("utf-8").split("data: ", 1)[1]))
|
||||
return events
|
||||
|
||||
events = asyncio.run(collect_events())
|
||||
|
||||
errors = [str(event.get("message", "")) for event in events if event.get("stage") == "agent"]
|
||||
self.assertEqual(agent.responses, [])
|
||||
self.assertTrue(any("连续两次未返回完整的 CDSL 工具 JSON" in error for error in errors))
|
||||
self.assertFalse(any("safety limit" in error for error in errors))
|
||||
diagnostics = sorted(settings.conversation_root.glob("conv_*/diagnostics/tool_call_*.json"))
|
||||
self.assertEqual(len(diagnostics), 2)
|
||||
records = [json.loads(path.read_text(encoding="utf-8")) for path in diagnostics]
|
||||
self.assertEqual([record["arguments"] for record in records], ['{"cdsl":', '{"cdsl":'])
|
||||
self.assertTrue(all(record["parse_error"] == "arguments are not valid JSON" for record in records))
|
||||
self.assertTrue(all(record["finish_reason"] == "length" for record in records))
|
||||
self.assertTrue(all(record["json_error"]["character"] == 8 for record in records))
|
||||
|
||||
def test_invalid_arguments_are_returned_to_the_model_for_retry(self) -> None:
|
||||
class RetryAgent(AgentService):
|
||||
def __init__(self, *args: object, **kwargs: object) -> None:
|
||||
super().__init__(*args, **kwargs)
|
||||
self.responses = [
|
||||
{
|
||||
"choices": [{"message": {
|
||||
"role": "assistant",
|
||||
"content": "I will search for a water cup reference.",
|
||||
"tool_calls": [{
|
||||
"id": "bad_call",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "search_cdsl_library",
|
||||
"arguments": '{"query":"water cup"}{"limit":3}',
|
||||
},
|
||||
}],
|
||||
}}],
|
||||
},
|
||||
{"choices": [{"message": {"role": "assistant", "content": "已修正工具参数。", "tool_calls": []}}]},
|
||||
]
|
||||
|
||||
async def _complete(self, *args: object, **kwargs: object) -> dict[str, object]:
|
||||
return self.responses.pop(0)
|
||||
|
||||
backend_root = Path(__file__).resolve().parents[1]
|
||||
with tempfile.TemporaryDirectory() as temporary_directory:
|
||||
temporary_root = Path(temporary_directory)
|
||||
provider = ProviderConfig("test", "Test", "https://example.invalid/v1", "test-key", (ProviderModel("test-model"),))
|
||||
settings = Settings(
|
||||
task_root=temporary_root / "tasks",
|
||||
conversation_root=temporary_root / "conversations",
|
||||
library_root=backend_root / "cdsl_library",
|
||||
engine_root=backend_root / "engine" / "cdsl_engine",
|
||||
llm_base_url=provider.base_url,
|
||||
llm_api_key=provider.api_key,
|
||||
llm_model="test-model",
|
||||
llm_timeout_s=1,
|
||||
default_provider_id="test",
|
||||
providers=(provider,),
|
||||
)
|
||||
store = WorkspaceStore(settings)
|
||||
agent = RetryAgent(settings, store, CdslLibrary(settings))
|
||||
message = ChatMessage(id="user_1", role="user", parts=[MessagePart(type="text", text="生成水杯")])
|
||||
|
||||
async def collect_events() -> list[dict[str, object]]:
|
||||
events: list[dict[str, object]] = []
|
||||
async for chunk in agent.stream([message], None, None):
|
||||
events.append(json.loads(chunk.decode("utf-8").split("data: ", 1)[1]))
|
||||
return events
|
||||
|
||||
events = asyncio.run(collect_events())
|
||||
|
||||
self.assertEqual(agent.responses, [])
|
||||
self.assertTrue(any(event.get("status") == "error" for event in events))
|
||||
self.assertFalse(any(event.get("stage") == "agent" for event in events))
|
||||
self.assertFalse(any("I will search" in str(event.get("text", "")) for event in events))
|
||||
self.assertTrue(any("已修正工具参数" in str(event.get("text", "")) for event in events))
|
||||
self.assertEqual(list(settings.task_root.glob("cad_*")), [])
|
||||
|
||||
def test_incomplete_cdsl_is_returned_to_the_model_without_creating_a_task(self) -> None:
|
||||
class RetryAgent(AgentService):
|
||||
def __init__(self, *args: object, **kwargs: object) -> None:
|
||||
super().__init__(*args, **kwargs)
|
||||
self.seen_messages: list[list[dict[str, object]]] = []
|
||||
self.required_tools: list[str | None] = []
|
||||
self.responses = [
|
||||
{
|
||||
"choices": [{"message": {
|
||||
"role": "assistant",
|
||||
"content": "",
|
||||
"tool_calls": [{
|
||||
"id": "incomplete_cdsl",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "generate_cdsl_model",
|
||||
"arguments": json.dumps({
|
||||
"cdsl": {"schema": "cad.cdsl.llm.v1"},
|
||||
"summary": "incomplete",
|
||||
}),
|
||||
},
|
||||
}],
|
||||
}}],
|
||||
},
|
||||
{
|
||||
"choices": [{"message": {
|
||||
"role": "assistant",
|
||||
"content": "",
|
||||
"tool_calls": [{
|
||||
"id": "corrected_cdsl",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "generate_cdsl_model",
|
||||
"arguments": json.dumps({
|
||||
"cdsl": {
|
||||
"schema": "cad.cdsl.llm.v1",
|
||||
"features": [{"id": "f01", "atomic_id": "extrude_add_blind", "depends_on": [], "params": {}, "sketch_id": "s01"}],
|
||||
"geometry": {"sketches": [{"id": "s01", "workplane": {}, "profile": {"type": "circle"}}]},
|
||||
},
|
||||
"summary": "complete",
|
||||
}),
|
||||
},
|
||||
}],
|
||||
}}],
|
||||
},
|
||||
{"choices": [{"message": {"role": "assistant", "content": "已补全模型。", "tool_calls": []}}]},
|
||||
]
|
||||
|
||||
async def _complete(self, messages: list[dict[str, object]], *args: object, **kwargs: object) -> dict[str, object]:
|
||||
self.seen_messages.append([dict(message) for message in messages])
|
||||
self.required_tools.append(kwargs.get("required_tool_name") if "required_tool_name" in kwargs else (args[3] if len(args) > 3 else None))
|
||||
return self.responses.pop(0)
|
||||
|
||||
async def _run_tool(self, name: str, arguments: dict[str, object], *args: object, **kwargs: object) -> tuple[dict[str, object], dict[str, object] | None]:
|
||||
if name == "generate_cdsl_model" and arguments.get("cdsl", {}).get("features"):
|
||||
return {"ok": True, "summary": "complete"}, None
|
||||
return await super()._run_tool(name, arguments, *args, **kwargs)
|
||||
|
||||
backend_root = Path(__file__).resolve().parents[1]
|
||||
with tempfile.TemporaryDirectory() as temporary_directory:
|
||||
temporary_root = Path(temporary_directory)
|
||||
provider = ProviderConfig("test", "Test", "https://example.invalid/v1", "test-key", (ProviderModel("test-model"),))
|
||||
settings = Settings(
|
||||
task_root=temporary_root / "tasks",
|
||||
conversation_root=temporary_root / "conversations",
|
||||
library_root=backend_root / "cdsl_library",
|
||||
engine_root=backend_root / "engine" / "cdsl_engine",
|
||||
llm_base_url=provider.base_url,
|
||||
llm_api_key=provider.api_key,
|
||||
llm_model="test-model",
|
||||
llm_timeout_s=1,
|
||||
default_provider_id="test",
|
||||
providers=(provider,),
|
||||
)
|
||||
agent = RetryAgent(settings, WorkspaceStore(settings), CdslLibrary(settings))
|
||||
message = ChatMessage(id="user_1", role="user", parts=[MessagePart(type="text", text="生成零件")])
|
||||
|
||||
async def collect_events() -> None:
|
||||
async for _ in agent.stream([message], None, None):
|
||||
pass
|
||||
|
||||
asyncio.run(collect_events())
|
||||
|
||||
tool_result = agent.seen_messages[1][-1]
|
||||
self.assertEqual(tool_result["role"], "tool")
|
||||
self.assertEqual(json.loads(str(tool_result["content"]))["code"], "INVALID_CDSL")
|
||||
self.assertEqual(agent.required_tools, [None, "generate_cdsl_model", None])
|
||||
self.assertEqual(list(settings.task_root.glob("cad_*")), [])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user