Files
cdsl-cad/backend/tests/test_author_guidance.py
T

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()