159 lines
6.8 KiB
Python
159 lines
6.8 KiB
Python
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import json
|
|
from pathlib import Path
|
|
import sys
|
|
import tempfile
|
|
import unittest
|
|
|
|
|
|
ROOT = Path(__file__).resolve().parents[2]
|
|
sys.path.insert(0, str(ROOT / "backend"))
|
|
|
|
from app.cad_agent.adapters.author_guidance import FileAuthorGuidance # noqa: E402
|
|
from app.cad_agent.application.workflow import ModelIdentity, WorkflowConfig, WorkflowCoordinator # noqa: E402
|
|
from app.cad_agent.domain.errors import ErrorCode, WorkflowError # noqa: E402
|
|
from app.cad_agent.domain.state import TaskPhase, TaskState # noqa: E402
|
|
|
|
|
|
GUIDANCE_ROOT = ROOT / "backend" / "agent" / "skills" / "cdsl-author-guidance"
|
|
PROFILE = ROOT / "backend" / "engine" / "cdsl_engine" / "profile_schema.json"
|
|
|
|
|
|
def atomic_ids() -> tuple[str, ...]:
|
|
return tuple(json.loads(PROFILE.read_text(encoding="utf-8"))["operation_contracts"])
|
|
|
|
|
|
class _Repository:
|
|
def __init__(self, state: TaskState) -> None:
|
|
self.state = state
|
|
self.usage_records: list[dict] = []
|
|
|
|
def get_state(self, _task_id: str) -> TaskState:
|
|
return self.state
|
|
|
|
def ledger_events(self, _task_id: str) -> list[dict]:
|
|
return []
|
|
|
|
def record_usage(self, _task_id: str, payload: dict) -> None:
|
|
self.usage_records.append(payload)
|
|
|
|
def record_tool_audit(self, _task_id: str, _payload: dict) -> None:
|
|
pass
|
|
|
|
|
|
class _Artifacts:
|
|
def read_source_requirements(self, _task_id: str) -> str:
|
|
return "Create a symmetric mounting plate."
|
|
|
|
def read_json(self, *_args: object) -> None:
|
|
return None
|
|
|
|
|
|
class _Runtime:
|
|
def supported_atomic_ids(self) -> tuple[str, ...]:
|
|
return atomic_ids()
|
|
|
|
|
|
class AuthorGuidanceTests(unittest.TestCase):
|
|
def test_manifest_covers_every_runtime_atomic_and_keeps_coordinate_core_at_minimum_budget(self) -> None:
|
|
guidance = FileAuthorGuidance(GUIDANCE_ROOT, max_chars=1_200)
|
|
covered: set[str] = set()
|
|
for atomic_id in atomic_ids():
|
|
selection = guidance.select(
|
|
phase=TaskPhase.FEATURE_PENDING,
|
|
atomic_id=atomic_id,
|
|
repair_required=False,
|
|
supported_atomic_ids=atomic_ids(),
|
|
)
|
|
self.assertTrue(selection.enabled, selection.fallback_reason)
|
|
self.assertIn("00-author-contract", selection.section_ids)
|
|
self.assertIn("03-coordinate-system-and-datums", selection.section_ids)
|
|
self.assertLessEqual(len(selection.content), 1_200)
|
|
self.assertIn("世界坐标", selection.content)
|
|
covered.update(section_id for section_id in selection.section_ids if section_id.startswith("op-"))
|
|
self.assertEqual(covered, {"op-extrude-add", "op-extrude-cut", "op-loft", "op-revolve", "op-hole", "op-reference", "op-pattern", "op-finish", "op-sphere", "op-primitives", "op-thread", "op-bend", "op-gear"})
|
|
|
|
def test_phase_repair_and_budget_selection_are_stable(self) -> None:
|
|
guidance = FileAuthorGuidance(GUIDANCE_ROOT, max_chars=3_600)
|
|
planning = guidance.select(
|
|
phase=TaskPhase.COMPILING_FEATURE_PLAN,
|
|
atomic_id="",
|
|
repair_required=False,
|
|
supported_atomic_ids=atomic_ids(),
|
|
)
|
|
repair = guidance.select(
|
|
phase=TaskPhase.AWAITING_ACTION,
|
|
atomic_id="fillet",
|
|
repair_required=True,
|
|
supported_atomic_ids=atomic_ids(),
|
|
)
|
|
self.assertEqual(planning.section_ids[:2], ("00-author-contract", "03-coordinate-system-and-datums"))
|
|
self.assertIn("02-parameters-and-derived-dimensions", planning.section_ids)
|
|
self.assertEqual(repair.section_ids[:3], ("00-author-contract", "03-coordinate-system-and-datums", "op-finish"))
|
|
self.assertIn("10-repair-and-best-effort", repair.section_ids)
|
|
|
|
def test_disabled_missing_and_invalid_corpus_fall_back_without_authoring_failure(self) -> None:
|
|
common = {
|
|
"phase": TaskPhase.FEATURE_PENDING,
|
|
"atomic_id": "extrude_add_blind",
|
|
"repair_required": False,
|
|
"supported_atomic_ids": atomic_ids(),
|
|
}
|
|
self.assertEqual(FileAuthorGuidance(GUIDANCE_ROOT, enabled=False).select(**common).fallback_reason, "guidance_disabled")
|
|
with tempfile.TemporaryDirectory() as temporary:
|
|
root = Path(temporary)
|
|
self.assertEqual(FileAuthorGuidance(root).select(**common).fallback_reason, "guidance_load_failed:FileNotFoundError")
|
|
(root / "manifest.json").write_text("{}", encoding="utf-8")
|
|
self.assertEqual(FileAuthorGuidance(root).select(**common).fallback_reason, "guidance_load_failed:ValueError")
|
|
|
|
def test_author_context_receives_guidance_but_keeps_the_existing_tool_instruction(self) -> None:
|
|
state = TaskState("cad_123456abcdef", TaskPhase.DRAFTING_REQUIREMENTS_DOCUMENT, 1)
|
|
workflow = WorkflowCoordinator(
|
|
WorkflowConfig(max_turns=8, format_error_limit=2),
|
|
_Repository(state),
|
|
_Artifacts(),
|
|
_Runtime(),
|
|
object(),
|
|
object(),
|
|
object(),
|
|
object(),
|
|
FileAuthorGuidance(GUIDANCE_ROOT),
|
|
)
|
|
messages, selection = workflow._author_context(state.task_id, [])
|
|
system = str(messages[0]["content"])
|
|
self.assertTrue(selection.enabled)
|
|
self.assertIn("Use exactly one offered structured tool call", system)
|
|
self.assertIn("Coordinate System And Datums", system)
|
|
self.assertIn("世界坐标", system)
|
|
|
|
def test_invalid_author_tool_call_retains_guidance_usage_metadata(self) -> None:
|
|
state = TaskState("cad_123456abcdef", TaskPhase.DRAFTING_REQUIREMENTS_DOCUMENT, 1)
|
|
repository = _Repository(state)
|
|
class _Models:
|
|
async def call_tool(self, **_kwargs: object) -> dict:
|
|
return {"tool_calls": [], "usage": {"prompt_tokens": 3, "completion_tokens": 1, "total_tokens": 4}}
|
|
workflow = WorkflowCoordinator(
|
|
WorkflowConfig(max_turns=8, format_error_limit=2),
|
|
repository,
|
|
_Artifacts(),
|
|
_Runtime(),
|
|
_Models(),
|
|
object(),
|
|
object(),
|
|
object(),
|
|
FileAuthorGuidance(GUIDANCE_ROOT),
|
|
)
|
|
tool = {"type": "function", "function": {"name": "write_requirements_document", "parameters": {"type": "object"}}}
|
|
result = asyncio.run(workflow._author_turn(state.task_id, ModelIdentity("provider", "model"), [tool], []))
|
|
self.assertIsInstance(result, WorkflowError)
|
|
self.assertEqual(result.code, ErrorCode.AUTHOR_FORMAT_INVALID)
|
|
self.assertEqual(repository.usage_records[0]["guidance_enabled"], True)
|
|
self.assertIn("03-coordinate-system-and-datums", repository.usage_records[0]["guidance_section_ids"])
|
|
self.assertEqual(repository.usage_records[0]["retry_reason"], "invalid_tool_call")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|