Files
cdsl-cad/backend/app/services/incremental_generation.py
T
2026-08-26 14:13:11 +08:00

510 lines
30 KiB
Python

"""Persistent, full-rebuild orchestration for node-by-node CDSL authoring."""
from __future__ import annotations
import asyncio
from collections.abc import AsyncIterator, Awaitable, Callable
from copy import deepcopy
import json
from pathlib import Path
import secrets
from typing import Any
from app.services.cdsl_fragment import CdslFragmentError, cdsl_sha256, materialize_fragment, validate_fragment
from app.services.engine_service import QualityVerificationError, build_revision, load_engine, normalize_cdsl_for_engine, validate_cdsl
from app.services.generation_plan import GenerationPlanError, descendant_closure, mark_nodes_stale, validate_generation_plan
from app.services.quality import validate_verification
from app.services.review_renderer import ReviewRenderError, render_checkpoint, renderer_status
from app.services.storage import WorkspaceStore, write_json
from app.services.visual_review import VisualReviewError, review_checkpoint
from app.settings import ProviderConfig, ProviderModel, Settings
Completion = Callable[[list[dict[str, Any]], list[dict[str, Any]], ProviderConfig, ProviderModel, str | None], Awaitable[dict[str, Any]]]
PLAN_TOOL = {
"type": "function",
"function": {
"name": "plan_generation_task",
"description": "Create the complete immutable requirement list and executable feature DAG before authoring any CDSL.",
"parameters": {
"type": "object",
"properties": {
"schema_version": {"type": "string", "const": "cad.generation-plan.v2"},
"plan_id": {"type": "string", "minLength": 1},
"requirements": {
"type": "array",
"minItems": 1,
"items": {
"type": "object",
"properties": {
"id": {"type": "string", "minLength": 1},
"source": {"enum": ["explicit", "assumption"]},
"priority": {"enum": ["hard", "soft"]},
"description": {"type": "string", "minLength": 1},
"value": {},
"unit": {"type": "string"},
"tolerance": {},
},
"required": ["id", "source", "priority", "description"],
"additionalProperties": False,
},
},
"assumptions": {"type": "array", "items": {"type": "string"}},
"nodes": {
"type": "array",
"minItems": 1,
"items": {
"type": "object",
"properties": {
"id": {"type": "string", "minLength": 1},
"intent": {"type": "string"},
"atomic_id": {"type": "string", "minLength": 1},
"depends_on": {"type": "array", "items": {"type": "string"}},
"requires_topology": {"type": "boolean"},
"topology_query": {"type": "object"},
"requirement_ids": {"type": "array", "items": {"type": "string"}},
"verification_rules": {"type": "array", "items": {"type": "object"}},
"review_targets": {"type": "array", "items": {"type": "object"}},
},
"required": ["id", "intent", "atomic_id", "depends_on", "requirement_ids", "verification_rules", "review_targets"],
"additionalProperties": False,
},
},
},
"required": ["schema_version", "plan_id", "requirements", "assumptions", "nodes"],
"additionalProperties": False,
},
},
}
FRAGMENT_TOOL = {
"type": "function",
"function": {
"name": "generate_cdsl_fragment",
"description": "Generate only the active plan node's additive CDSL fragment. Never replace or mutate existing CDSL.",
"parameters": {
"type": "object",
"properties": {
"schema_version": {"type": "string", "const": "cad.cdsl-fragment.v1"},
"node_id": {"type": "string", "minLength": 1},
"base_revision_id": {"type": "string"},
"base_cdsl_sha256": {"type": "string", "minLength": 64, "maxLength": 64},
"required_snapshot_id": {"type": "string"},
"add_sketches": {"type": "array", "items": {"type": "object"}},
"add_features": {"type": "array", "items": {"type": "object"}},
"verification_rules": {"type": "array", "items": {"type": "object"}},
"assumptions": {"type": "array", "items": {"type": "string"}},
},
"required": ["schema_version", "node_id", "base_revision_id", "base_cdsl_sha256", "add_sketches", "add_features", "verification_rules", "assumptions"],
"additionalProperties": False,
},
},
}
class IncrementalGenerationError(RuntimeError):
pass
def _tool_response(response: dict[str, Any], expected_name: str) -> dict[str, Any]:
try:
call = response["choices"][0]["message"]["tool_calls"][0]
if call["function"]["name"] != expected_name:
raise KeyError("wrong tool")
result = json.loads(call["function"]["arguments"])
except (KeyError, IndexError, TypeError, json.JSONDecodeError) as error:
raise IncrementalGenerationError(f"Author did not return a valid {expected_name} call") from error
if not isinstance(result, dict):
raise IncrementalGenerationError(f"{expected_name} arguments must be an object")
return result
def _node_by_id(spec: dict[str, Any], node_id: str) -> dict[str, Any]:
node = next((item for item in spec.get("nodes") or () if isinstance(item, dict) and item.get("id") == node_id), None)
if node is None:
raise IncrementalGenerationError(f"Generation plan has no node {node_id}")
return node
def _fragment_node_context(node: dict[str, Any]) -> dict[str, Any]:
"""Expose semantic node intent, not backend-owned CDSL implementation IDs."""
fields = (
"id", "intent", "atomic_id", "depends_on", "requires_topology",
"requires_sketch", "topology_query", "requirement_ids",
"verification_rules", "review_targets",
)
return {field: deepcopy(node[field]) for field in fields if field in node}
def _ready_node(spec: dict[str, Any], completed: set[str], has_topology: bool) -> dict[str, Any] | None:
for node in spec.get("nodes") or ():
if not isinstance(node, dict) or node.get("status") == "completed":
continue
if all(str(item) in completed for item in node.get("depends_on") or ()) and (not node.get("requires_topology") or has_topology):
return node
return None
def _error_code(error: Exception) -> str:
message = str(error)
for code in (
"SELECTOR_AMBIGUOUS", "SELECTOR_NOT_FOUND", "SELECTOR_GEOMETRY_MISMATCH",
"TOPOLOGY_SNAPSHOT_STALE", "TOPOLOGY_REQUIRED", "VERIFICATION_FAILED", "SELECTOR_OWNER_REQUIRED",
):
if code in message:
return code
return type(error).__name__.upper()
def _mark_affected_nodes_stale(plan: dict[str, Any], node_ids: list[str], *, reason: str) -> tuple[dict[str, Any], set[str]]:
"""Invalidate the union of every affected node's downstream closure."""
updated = deepcopy(plan)
stale: set[str] = set()
for node_id in dict.fromkeys(node_ids):
stale.update(descendant_closure(updated, node_id))
updated = mark_nodes_stale(updated, node_id, reason=reason)
return updated, stale
def _source_image_paths(store: WorkspaceStore, conversation: dict[str, Any]) -> list[Path]:
conversation_id = str(conversation.get("conversation_id") or "")
paths: list[Path] = []
for attachment in conversation.get("attachments") or ():
if not isinstance(attachment, dict) or attachment.get("kind") != "image" or not conversation_id:
continue
try:
paths.append(store.conversation_attachment_path(conversation_id, str(attachment.get("path") or "")))
except ValueError:
continue
return paths
class IncrementalGenerationRunner:
"""The agent-facing controller. It is deliberately full-rebuild and restart-safe."""
def __init__(self, settings: Settings, store: WorkspaceStore, complete: Completion) -> None:
self.settings = settings
self.store = store
self._complete = complete
async def _call(self, messages: list[dict[str, Any]], provider: ProviderConfig, model: ProviderModel, tool: dict[str, Any], name: str) -> dict[str, Any]:
response = await self._complete(messages, [tool], provider, model, name)
return _tool_response(response, name)
async def run(
self,
*,
task_id: str,
request: str,
conversation: dict[str, Any],
provider: ProviderConfig,
model: ProviderModel,
author_messages: list[dict[str, Any]],
part_skills: dict[str, Any] | None = None,
references: list[str] | None = None,
already_started: bool = False,
) -> AsyncIterator[tuple[str, dict[str, Any]]]:
plan_diagnostic_path = ""
try:
task = self.store.ensure_task(task_id or None, request)
task_id = str(task["task_id"])
# Configuration is a start gate: visual review is required, never silently skipped.
self.settings.resolve_review_model()
renderer_ready, renderer_error = renderer_status()
if not renderer_ready:
raise IncrementalGenerationError(renderer_error)
task = self.store.read_task(task_id) if already_started else self.store.start_generation(task_id, request=request)
if not isinstance(task, dict):
raise IncrementalGenerationError("Generation task is unavailable")
yield "generation_plan", {"taskId": task_id, "status": "running"}
engine = load_engine(self.settings)
persisted_spec = self.store.read_generation_spec(task_id)
if persisted_spec is not None:
spec = validate_generation_plan(
persisted_spec,
supported_atomic_ids=getattr(engine, "SUPPORTED_ATOMIC_IDS", ()),
task_id=task_id,
)
yield "generation_plan", {
"taskId": task_id, "status": "success", "planId": spec["plan_id"], "resumed": True,
"requirements": spec["requirements"],
"nodes": [{"id": node["id"], "intent": node["intent"], "status": node.get("status", "planned")} for node in spec["nodes"]],
}
else:
planning_messages = [
{
"role": "system",
"content": (
"Create one complete cad.generation-plan.v2 before creating geometry. "
"Turn every user constraint into a requirement with source explicit or assumption; "
"use source assumption for missing dimensions and never ask the user questions. "
"Every hard requirement must belong to at least one node. Use only runtime-supported atomic ids. "
"Do not output expected_feature_ids or expected_sketch_ids: the backend derives all CDSL object ids from node.id."
),
},
*author_messages,
]
raw_spec = await self._call(planning_messages, provider, model, PLAN_TOOL, "plan_generation_task")
try:
spec = validate_generation_plan(
raw_spec,
supported_atomic_ids=getattr(engine, "SUPPORTED_ATOMIC_IDS", ()),
task_id=task_id,
)
except GenerationPlanError as error:
plan_diagnostic_path = self.store.write_generation_failure(task_id, {
"schema_version": "cad.generation-plan-diagnostic.v1",
"stage": "generation_plan_validation",
"message": str(error),
"raw_plan": raw_spec,
})
raise
self.store.write_generation_spec(task_id, spec)
yield "generation_plan", {
"taskId": task_id, "status": "success", "planId": spec["plan_id"],
"requirements": spec["requirements"], "nodes": [{"id": node["id"], "intent": node["intent"], "status": "planned"} for node in spec["nodes"]],
}
completed: set[str] = {
str(node["id"]) for node in spec["nodes"] if node.get("status") == "completed"
}
last_built: dict[str, Any] | None = None
task = self.store.read_task(task_id) or task
active_revision_id = str(task.get("active_revision") or "")
# A process can stop between build and review. That checkpoint is
# not a legal base revision, so recover its parent before resuming.
active_record = next(
(item for item in task.get("revisions") or () if isinstance(item, dict) and item.get("revision_id") == active_revision_id),
None,
)
if isinstance(active_record, dict) and active_record.get("visibility") == "checkpoint":
current_node = str(active_record.get("node_id") or "")
if current_node and current_node not in completed:
recovered = str(active_record.get("parent_revision_id") or "")
self.store.rollback_to_revision(task_id, recovered, branch_id=f"branch_{secrets.token_hex(4)}")
active_revision_id = recovered
task = self.store.read_task(task_id) or task
base_path = self.store.current_cdsl_path(task_id)
base_cdsl = json.loads(base_path.read_text(encoding="utf-8")) if base_path and base_path.is_file() else None
while True:
topology_path = self.store.current_topology_path(task_id)
topology = json.loads(topology_path.read_text(encoding="utf-8")) if topology_path and topology_path.is_file() else None
node = _ready_node(spec, completed, bool(topology and topology.get("records")))
if node is None:
if len(completed) == len(spec["nodes"]):
self.store.finish_generation(task_id, lifecycle="completed")
if last_built is not None:
yield "cad_result", self._result_payload(last_built, lifecycle="completed", checkpoint=False)
yield "task_terminal", {"taskId": task_id, "lifecycle": "completed", "revisionId": str((self.store.read_task(task_id) or {}).get("published_revision") or "")}
return
waiting = [item["id"] for item in spec["nodes"] if item.get("id") not in completed]
raise IncrementalGenerationError("No executable plan node is ready: " + ", ".join(waiting))
node_id = str(node["id"])
self.store.set_active_node(task_id, node_id)
yield "checkpoint", {"taskId": task_id, "nodeId": node_id, "status": "authoring"}
attempts = node.setdefault("attempts", {"authoring": 0, "repair": 0, "replan": 0})
feedback = ""
while True:
attempt_kind = "authoring" if int(attempts.get("authoring") or 0) < self.settings.node_authoring_attempts else "repair"
if attempt_kind == "repair" and int(attempts.get("repair") or 0) >= self.settings.node_repair_attempts:
if int(attempts.get("replan") or 0) >= self.settings.node_replan_attempts:
raise IncrementalGenerationError(f"Node {node_id} exhausted its authoring, repair, and replan budgets: {feedback}")
attempts["replan"] = int(attempts.get("replan") or 0) + 1
spec = await self._replan(spec, node_id, feedback, provider, model, author_messages, engine)
self.store.write_generation_spec(task_id, spec)
node = _node_by_id(spec, node_id)
node["attempts"] = {"authoring": 0, "repair": 0, "replan": attempts["replan"]}
attempts = node["attempts"]
yield "rollback", {"taskId": task_id, "nodeId": node_id, "reason": "node_replan"}
continue
attempts[attempt_kind] = int(attempts.get(attempt_kind) or 0) + 1
required_snapshot = str((topology or {}).get("snapshot_id") or "") if node.get("requires_topology") else ""
node_requirement_ids = set(node.get("requirement_ids") or ())
author_context = {
"active_node": _fragment_node_context(node),
"requirements": [
requirement for requirement in spec["requirements"]
if requirement.get("id") in node_requirement_ids
],
"base_revision_id": active_revision_id,
"base_cdsl_sha256": cdsl_sha256(base_cdsl),
"required_snapshot_id": required_snapshot,
"base_cdsl": base_cdsl,
"topology": topology if required_snapshot else None,
"previous_failure": feedback,
}
fragment_messages = [
{
"role": "system",
"content": (
"Author exactly one additive CDSL fragment for the active node. Do not change existing CDSL or generate unsupported topology selectors. "
"Do not set output id, sketch_id, or depends_on: the backend assigns them. "
"Use owner_node_id instead of owner_feature_id and source_node_ids instead of source_feature_ids when referring to plan nodes."
),
},
*author_messages,
{"role": "user", "content": json.dumps(author_context, ensure_ascii=False)},
]
try:
raw_fragment = await self._call(fragment_messages, provider, model, FRAGMENT_TOOL, "generate_cdsl_fragment")
fragment = validate_fragment(
raw_fragment, plan=spec, node_id=node_id, base_revision_id=active_revision_id,
base_cdsl=base_cdsl, required_snapshot_id=required_snapshot,
)
cdsl = materialize_fragment(base_cdsl, fragment)
cdsl, repairs = normalize_cdsl_for_engine(cdsl)
validate_cdsl(cdsl, engine)
rules = [*node.get("verification_rules", []), *fragment.get("verification_rules", [])]
verification = {"rules": rules} if rules else None
validate_verification(verification, cdsl)
fragment_base_revision = str(fragment["base_revision_id"])
built = await asyncio.to_thread(
build_revision,
settings=self.settings, store=self.store, task_id=task_id, request=request, cdsl=cdsl,
reference_ids=references or [], summary=node.get("intent") or node_id,
parent_revision_id=active_revision_id or None,
operation={"type": "cdsl_fragment", "node_id": node_id},
part_skills=part_skills, generation_assumptions=[*spec.get("assumptions", []), *fragment.get("assumptions", [])],
verification=verification, node_id=node_id, fragment=fragment,
branch_id=str((self.store.read_task(task_id) or {}).get("active_branch_id") or "main"), visibility="checkpoint",
)
candidate_revision_id = str(built["revision_id"])
try:
render_dir = self.store.revision_dir(task_id, candidate_revision_id) / "review"
manifest = await asyncio.to_thread(
render_checkpoint,
self.settings,
step_path=self.store.artifact_path(task_id, str(built["step_path"])),
output_dir=render_dir,
review_targets=node.get("review_targets"),
)
final_checkpoint = len(completed) + 1 == len(spec["nodes"])
review_requirements = spec["requirements"] if final_checkpoint else [
requirement for requirement in spec["requirements"]
if requirement.get("id") in set(node.get("requirement_ids") or ())
]
review = await review_checkpoint(
self.settings, manifest=manifest, requirements=review_requirements, node_id=node_id,
deterministic_report={"quality_status": built.get("quality_status"), "verification": built.get("verification_summary", {})},
source_images=_source_image_paths(self.store, conversation) if not completed or final_checkpoint else [],
final_checkpoint=final_checkpoint,
)
except (ReviewRenderError, VisualReviewError) as error:
self.store.rollback_to_revision(task_id, fragment_base_revision, branch_id=f"branch_{secrets.token_hex(4)}")
active_revision_id = fragment_base_revision
rollback_path = self.store.current_cdsl_path(task_id)
base_cdsl = json.loads(rollback_path.read_text(encoding="utf-8")) if rollback_path and rollback_path.is_file() else None
raise error
active_revision_id = candidate_revision_id
base_cdsl = cdsl
manifest_relative = (render_dir / "render-manifest.json").relative_to(self.store.task_dir(task_id)).as_posix()
review_relative = (render_dir / "visual-review.json").relative_to(self.store.task_dir(task_id)).as_posix()
write_json(render_dir / "visual-review.json", review)
self.store.update_revision_metadata(task_id, active_revision_id, {"render_manifest_path": manifest_relative, "visual_review_path": review_relative})
yield "render_review", {"taskId": task_id, "revisionId": active_revision_id, "nodeId": node_id, "review": review}
if review["verdict"] == "repair" and float(review["confidence"]) >= 0.85:
affected_nodes = [
str(item) for item in review.get("affected_node_ids") or ()
if any(str(candidate.get("id") or "") == str(item) for candidate in spec.get("nodes") or ())
] or [node_id]
rollback_base = self.store.rollback_anchor_for_nodes(
task_id,
affected_nodes,
fallback_revision_id=fragment_base_revision,
)
branch = f"branch_{secrets.token_hex(4)}"
self.store.rollback_to_revision(task_id, rollback_base, branch_id=branch)
spec, stale_nodes = _mark_affected_nodes_stale(
spec, affected_nodes, reason="high_confidence_visual_review",
)
self.store.write_generation_spec(task_id, spec)
completed.difference_update(stale_nodes)
active_revision_id = rollback_base
rollback_path = self.store.current_cdsl_path(task_id)
base_cdsl = json.loads(rollback_path.read_text(encoding="utf-8")) if rollback_path and rollback_path.is_file() else None
feedback = "High-confidence visual review requires correction: " + "; ".join(review.get("evidence") or [])
yield "rollback", {
"taskId": task_id, "nodeId": node_id, "revisionId": active_revision_id,
"reason": "visual_review", "affectedNodeIds": affected_nodes,
}
continue
completed.add(node_id)
node["status"] = "completed"
node.pop("stale", None)
self.store.write_generation_spec(task_id, spec)
last_built = built
payload = self._result_payload(built, lifecycle="running", checkpoint=True)
yield "checkpoint", {"taskId": task_id, "nodeId": node_id, "status": "success", "revisionId": active_revision_id}
yield "cad_result", payload
break
except (CdslFragmentError, GenerationPlanError, QualityVerificationError, ReviewRenderError, VisualReviewError, ValueError, RuntimeError) as error:
feedback = str(error)
failure = {
"schema_version": "cad.generation-failure.v1",
"node_id": node_id,
"stage": attempt_kind,
"error_code": _error_code(error),
"message": feedback,
"requirement_ids": list(node.get("requirement_ids") or ()),
"selector": {"required_snapshot_id": required_snapshot},
"geometry_delta": {},
"recommended_rollback_revision": fragment_base_revision if "fragment_base_revision" in locals() else active_revision_id,
}
failure_path = self.store.write_generation_failure(task_id, failure)
yield "checkpoint", {"taskId": task_id, "nodeId": node_id, "status": "error", "attempt": attempt_kind, "message": feedback}
# A failed build never becomes the active base; retries are safe and deterministic.
continue
except Exception as error:
failure = {
"schema_version": "cad.generation-failure.v1",
"message": str(error),
"active_node_id": str((self.store.read_task(task_id) or {}).get("active_node_id") or ""),
}
if plan_diagnostic_path:
failure["plan_diagnostic_path"] = plan_diagnostic_path
self.store.finish_generation(task_id, lifecycle="failed", failure=failure)
yield "task_terminal", {"taskId": task_id, "lifecycle": "failed", "message": str(error)}
async def _replan(
self,
spec: dict[str, Any],
node_id: str,
feedback: str,
provider: ProviderConfig,
model: ProviderModel,
author_messages: list[dict[str, Any]],
engine: Any,
) -> dict[str, Any]:
raw = await self._call([
{"role": "system", "content": "Replan only the failed node and its downstream nodes. Preserve completed node definitions and all requirements."},
*author_messages,
{"role": "user", "content": json.dumps({"existing_plan": spec, "failed_node_id": node_id, "failure": feedback}, ensure_ascii=False)},
], provider, model, PLAN_TOOL, "plan_generation_task")
next_spec = validate_generation_plan(raw, supported_atomic_ids=getattr(engine, "SUPPORTED_ATOMIC_IDS", ()), task_id=str(spec.get("task_id") or ""))
stale = {item["id"] for item in mark_nodes_stale(spec, node_id, reason="replan").get("nodes") or [] if item.get("stale")}
prior = {item["id"]: item for item in spec.get("nodes") or []}
for node in next_spec["nodes"]:
if node["id"] not in stale and node["id"] in prior:
old = prior[node["id"]]
for key in ("atomic_id", "depends_on", "cdsl_feature_ids"):
if node.get(key) != old.get(key):
raise IncrementalGenerationError(f"Replan changed non-stale node {node['id']}")
return next_spec
@staticmethod
def _result_payload(built: dict[str, Any], *, lifecycle: str, checkpoint: bool) -> dict[str, Any]:
return {
"taskId": built["task_id"], "revisionId": built["revision_id"],
"cdslPath": built["cdsl_path"], "stepPath": built["step_path"], "glbPath": built["glb_path"],
"reportPath": built["report_path"], "parametersPath": built.get("parameters_path"),
"selectorPath": built.get("selector_path"), "edgesPath": built.get("edges_path"), "topologyPath": built.get("topology_path"),
"summary": built.get("summary") or "CDSL checkpoint", "referenceIds": built.get("reference_ids") or [],
"engine": built.get("engine") or "cdsl_only", "qualityStatus": built.get("quality_status") or "",
"qualityPath": built.get("quality_path"), "assumptions": built.get("generation_assumptions") or [],
"snapshotPaths": built.get("snapshot_paths") or [], "snapshotStatus": built.get("snapshot_status") or "unavailable",
"lifecycle": lifecycle, "checkpoint": checkpoint,
}