225 lines
13 KiB
Python
225 lines
13 KiB
Python
"""HTTP/SSE delivery adapter for the single-stage CAD protocol."""
|
|
from __future__ import annotations
|
|
|
|
import asyncio
|
|
from collections.abc import AsyncIterator
|
|
from hashlib import sha256
|
|
import re
|
|
import secrets
|
|
from typing import Any
|
|
|
|
from app.cad_agent.application.workflow import ModelIdentity
|
|
from app.cad_agent.composition import CadServices, compose_cad_services
|
|
from app.cad_agent.domain.errors import ErrorCode
|
|
from app.cad_agent.domain.state import TaskPhase, transition
|
|
from app.models.contracts import ChatMessage
|
|
from app.services.library import CdslLibrary
|
|
from app.services.sse import event
|
|
from app.services.storage import WorkspaceStore, now_iso
|
|
from app.settings import Settings
|
|
|
|
|
|
def text_from_message(message: ChatMessage) -> str:
|
|
return "\n".join(part.text or "" for part in message.parts if part.type == "text").strip()
|
|
|
|
|
|
_EVENT_LABELS = {
|
|
"requirements_ready": "需求分析",
|
|
"authoring_cdsl_ready": "完整 CDSL",
|
|
"cdsl_compiled": "CDSL 编译",
|
|
"build_result": "CAD 构建",
|
|
"repair_started": "CDSL 修复",
|
|
"task_terminal": "生成任务",
|
|
}
|
|
|
|
|
|
def _progress(name: str, payload: dict[str, Any]) -> dict[str, Any]:
|
|
lifecycle = str(payload.get("lifecycle") or "")
|
|
status = "error" if lifecycle == "failed" or payload.get("status") in {"failed", "repair_required"} else "waiting" if lifecycle == "waiting_for_user" else "success" if lifecycle == "completed" else "running"
|
|
return {**payload, "step": name, "label": _EVENT_LABELS.get(name, name), "status": status}
|
|
|
|
|
|
class AgentService:
|
|
def __init__(self, settings: Settings, store: WorkspaceStore, library: CdslLibrary) -> None:
|
|
self.settings, self.store, self.library = settings, store, library
|
|
self.cad: CadServices = compose_cad_services(settings)
|
|
self._autonomous_runs: dict[str, asyncio.Task[None]] = {}
|
|
|
|
async def resume_running_tasks(self) -> None:
|
|
if not self.settings.resume_running_tasks_on_startup:
|
|
return
|
|
try:
|
|
provider, model = self.settings.resolve_model(None, None)
|
|
except ValueError:
|
|
return
|
|
author = ModelIdentity(provider.id, model.id)
|
|
for task_id in self.cad.repository.running_task_ids():
|
|
if task_id not in self._autonomous_runs:
|
|
self._autonomous_runs[task_id] = asyncio.create_task(self._consume_discarding(task_id, author), name=f"resume-cad-single-stage-{task_id}")
|
|
|
|
async def _consume_discarding(self, task_id: str, author: ModelIdentity) -> None:
|
|
try:
|
|
async for _name, _payload in self.cad.workflow.run(task_id=task_id, author=author):
|
|
await self.cad.outbox.dispatch_pending(task_id=task_id)
|
|
finally:
|
|
await self.cad.outbox.dispatch_pending(task_id=task_id)
|
|
self._autonomous_runs.pop(task_id, None)
|
|
|
|
async def cancel(self, task_id: str) -> dict[str, Any] | None:
|
|
state = self.cad.repository.get_state(task_id)
|
|
if state is None:
|
|
return None
|
|
if state.phase not in {TaskPhase.COMPLETED, TaskPhase.FAILED, TaskPhase.CANCELLED}:
|
|
cancelled = transition(state, "cancelled", error=ErrorCode.CANCELLED)
|
|
self.cad.repository.compare_and_swap(cancelled, events=[{"event": "task_cancelled", "active_revision": state.active_revision}])
|
|
await self.cad.outbox.dispatch_pending(task_id=task_id)
|
|
running = self._autonomous_runs.pop(task_id, None)
|
|
if running and not running.done():
|
|
running.cancel()
|
|
return self.cad.repository.get_task_projection(task_id)
|
|
|
|
async def resume_retry(self, task_id: str) -> dict[str, Any] | None:
|
|
state = self.cad.repository.get_state(task_id)
|
|
if state is None:
|
|
return None
|
|
if state.phase != TaskPhase.FAILED or state.retry_from_phase is None:
|
|
raise ValueError("Only a retryable failed CAD task can be resumed")
|
|
if task_id in self._autonomous_runs and not self._autonomous_runs[task_id].done():
|
|
raise ValueError("The CAD task is already running")
|
|
provider, model = self.settings.resolve_model(None, None)
|
|
if not self.cad.workflow.resume(task_id):
|
|
raise ValueError("The CAD task no longer has a retry checkpoint")
|
|
self._autonomous_runs[task_id] = asyncio.create_task(self._consume_discarding(task_id, ModelIdentity(provider.id, model.id)), name=f"resume-cad-single-stage-{task_id}")
|
|
return self.cad.repository.get_task_projection(task_id)
|
|
|
|
async def stream(self, messages: list[ChatMessage], conversation_id: str | None, selected_task_id: str | None, provider_id: str | None = None, model_id: str | None = None, viewer_context: list[dict[str, Any]] | None = None) -> AsyncIterator[bytes]:
|
|
del viewer_context
|
|
latest = next((item for item in reversed(messages) if item.role == "user"), None)
|
|
if latest is None or not text_from_message(latest):
|
|
yield event("cad_error", {"stage": "request", "message": "A non-empty user request is required."})
|
|
yield event("done", {})
|
|
return
|
|
request = text_from_message(latest)
|
|
conversation = self.store.ensure_conversation(conversation_id)
|
|
selected = str(selected_task_id or conversation.get("current_task_id") or "")
|
|
current = self.cad.repository.get_task_projection(selected) if selected else None
|
|
resumed = False
|
|
if current and current.get("lifecycle") == "waiting_for_user":
|
|
resumed = self.cad.workflow.resume_with_user_clarification(selected, request, message_id=latest.id)
|
|
if not resumed:
|
|
yield event("cad_error", {"stage": "request", "message": "The CAD task cannot apply this clarification."})
|
|
yield event("done", {})
|
|
return
|
|
if current and current.get("lifecycle") == "running":
|
|
yield event("cad_error", {"stage": "request", "message": "该 CAD 任务正在生成,完成后才能继续。"})
|
|
yield event("done", {})
|
|
return
|
|
task_id = selected if resumed else f"cad_{secrets.token_hex(6)}"
|
|
if not resumed:
|
|
try:
|
|
source_blocks, image_inputs = self._task_inputs(conversation, request)
|
|
self.cad.workflow.create_task(task_id, request, source_blocks=source_blocks, image_inputs=image_inputs)
|
|
except ValueError as error:
|
|
yield event("cad_error", {"stage": "request", "message": str(error)})
|
|
yield event("done", {})
|
|
return
|
|
self.store.append_conversation_message(conversation["conversation_id"], latest.model_dump(), task_id)
|
|
yield event("progress", _progress("task_started", {"taskId": task_id, "message": "CAD task started."}))
|
|
try:
|
|
provider, model = self.settings.resolve_model(provider_id, model_id)
|
|
except ValueError as error:
|
|
state = self.cad.repository.get_state(task_id)
|
|
if state:
|
|
failed = transition(state, "failed", error=ErrorCode.MODEL_STRUCTURED_OUTPUT_UNSUPPORTED)
|
|
self.cad.repository.compare_and_swap(failed, events=[{"event": "model_configuration_invalid", "message": str(error)}])
|
|
yield event("cad_error", {"stage": "configuration", "message": str(error)})
|
|
yield event("done", {})
|
|
return
|
|
author = ModelIdentity(provider.id, model.id)
|
|
queue: asyncio.Queue[tuple[str, dict[str, Any]] | None] = asyncio.Queue()
|
|
parts: list[dict[str, Any]] = []
|
|
|
|
async def consume() -> None:
|
|
try:
|
|
async for name, payload in self.cad.workflow.run(task_id=task_id, author=author):
|
|
decorated = {**payload, "taskId": task_id, "eventId": f"{task_id}_{secrets.token_hex(4)}", "timestamp": now_iso()}
|
|
parts.append({"type": "data-cad-progress", "id": decorated["eventId"], "data": _progress(name, decorated)})
|
|
await queue.put((name, decorated))
|
|
await self.cad.outbox.dispatch_pending(task_id=task_id)
|
|
if name == "task_terminal" and decorated.get("lifecycle") == "completed":
|
|
projection = self.cad.repository.get_task_projection(task_id) or {}
|
|
result = self._result_payload(task_id, projection)
|
|
if result is not None:
|
|
await queue.put(("cad_result", result))
|
|
finally:
|
|
self.store.append_conversation_message(conversation["conversation_id"], {"id": f"assistant_{secrets.token_hex(8)}", "role": "assistant", "parts": parts}, task_id)
|
|
self._autonomous_runs.pop(task_id, None)
|
|
await queue.put(None)
|
|
|
|
self._autonomous_runs[task_id] = asyncio.create_task(consume(), name=f"cad-single-stage-{task_id}")
|
|
while True:
|
|
item = await queue.get()
|
|
if item is None:
|
|
break
|
|
name, payload = item
|
|
if name == "task_terminal" and payload.get("lifecycle") == "failed":
|
|
yield event("cad_error", {"stage": "generation", "message": str(payload.get("message") or "CAD generation failed."), "code": payload.get("code")})
|
|
yield event(name, payload)
|
|
yield event("done", {})
|
|
|
|
@staticmethod
|
|
def _result_payload(task_id: str, projection: dict[str, Any]) -> dict[str, Any] | None:
|
|
revision_id = str(projection.get("published_revision") or projection.get("active_revision") or "")
|
|
revisions = projection.get("revisions") if isinstance(projection.get("revisions"), list) else []
|
|
revision = next((item for item in revisions if isinstance(item, dict) and item.get("revision_id") == revision_id), None)
|
|
if not isinstance(revision, dict):
|
|
return None
|
|
required = ("cdsl_path", "step_path", "glb_path", "report_path")
|
|
if not all(isinstance(revision.get(path), str) and revision[path] for path in required):
|
|
return None
|
|
return {
|
|
"taskId": task_id,
|
|
"revisionId": revision_id,
|
|
"cdslPath": revision["cdsl_path"],
|
|
"stepPath": revision["step_path"],
|
|
"glbPath": revision["glb_path"],
|
|
"reportPath": revision["report_path"],
|
|
"summary": str(revision.get("summary") or "CDSL CAD model"),
|
|
"referenceIds": list(revision.get("reference_ids") or []),
|
|
"engine": str(revision.get("engine") or "cdsl_only"),
|
|
"lifecycle": str(projection.get("lifecycle") or "completed"),
|
|
}
|
|
|
|
def _task_inputs(self, conversation: dict[str, Any], request: str) -> tuple[list[dict[str, Any]], list[dict[str, str]]]:
|
|
blocks = [{"text": paragraph} for paragraph in re.split(r"\n\s*\n", request) if paragraph.strip()]
|
|
image_inputs: list[dict[str, str]] = []
|
|
conversation_id = str(conversation.get("conversation_id") or "")
|
|
if not conversation_id:
|
|
raise ValueError("Conversation has no identifier")
|
|
for attachment in conversation.get("attachments") or ():
|
|
if not isinstance(attachment, dict):
|
|
continue
|
|
if str(attachment.get("conversation_id") or "") != conversation_id:
|
|
raise ValueError("Attachment does not belong to this conversation")
|
|
attachment_id, relative_path, digest = str(attachment.get("id") or ""), str(attachment.get("path") or ""), str(attachment.get("sha256") or "")
|
|
if not attachment_id or not relative_path or not re.fullmatch(r"[a-f0-9]{64}", digest):
|
|
raise ValueError("Attachment metadata is incomplete")
|
|
binary = self.store.conversation_attachment_path(conversation_id, relative_path)
|
|
if not binary.is_file() or sha256(binary.read_bytes()).hexdigest() != digest:
|
|
raise ValueError(f"Attachment is unavailable: {attachment.get('name') or attachment_id}")
|
|
kind = str(attachment.get("kind") or "")
|
|
if kind == "document":
|
|
text_path = self.store.conversation_attachment_path(conversation_id, str(attachment.get("extracted_path") or ""))
|
|
if not text_path.is_file():
|
|
raise ValueError(f"Attachment text is unavailable: {attachment.get('name') or attachment_id}")
|
|
text = text_path.read_text(encoding="utf-8", errors="replace").strip()
|
|
elif kind == "image":
|
|
text = f"Visual attachment {attachment.get('name') or attachment_id} (SHA-256 {digest})."
|
|
image_inputs.append({"path": str(binary), "mime": str(attachment.get("mime") or "image/*"), "sha256": digest})
|
|
else:
|
|
raise ValueError(f"Unsupported attachment kind: {kind or 'unknown'}")
|
|
if not text:
|
|
raise ValueError(f"Attachment source is empty: {attachment.get('name') or attachment_id}")
|
|
blocks.append({"text": text, "attachment": {"attachment_id": attachment_id, "name": str(attachment.get("name") or attachment_id), "kind": kind, "mime": str(attachment.get("mime") or "application/octet-stream"), "sha256": digest}})
|
|
return blocks, image_inputs
|