Files
cdsl-cad/backend/app/main.py
T
2026-08-19 19:34:30 +08:00

190 lines
7.6 KiB
Python

from __future__ import annotations
import json
from typing import Any
from fastapi import FastAPI, File, Form, HTTPException, UploadFile
from fastapi.responses import JSONResponse, StreamingResponse
from app.models.contracts import ChatRequest, ConversationPatch, ModifyRequest, ParameterUpdate
from app.services.engine_service import apply_parameter_updates, build_revision
from app.services.agent_service import AgentService
from app.services.library import CdslLibrary
from app.services.storage import WorkspaceStore, safe_conversation_id, safe_task_id
from app.services.attachments import attachment_record, classify_upload, extract_document_text
from app.settings import get_settings
settings = get_settings()
store = WorkspaceStore(settings)
library = CdslLibrary(settings)
agent = AgentService(settings, store, library)
app = FastAPI(title="CDSL CAD Agent API", version="0.1.0")
@app.get("/health")
async def health() -> dict[str, Any]:
return {
"ok": True,
"service": "cdsl-cad-backend",
"llm_configured": settings.llm_configured,
"library_index": (settings.library_root / "index" / "catalog.json").is_file(),
}
@app.get("/v1/config")
async def config() -> dict[str, Any]:
providers = []
for provider in settings.providers:
if not provider.configured:
continue
providers.append({
"id": provider.id,
"label": provider.label,
"models": [{"id": model.id, "vision": model.vision} for model in provider.models],
})
return {
"default_provider": settings.default_provider_id,
"default_model": settings.llm_model,
"providers": providers,
"model": settings.llm_model,
"configured": settings.llm_configured,
"library_samples": library.count(),
}
@app.post("/v1/chat/stream")
async def chat_stream(payload: ChatRequest) -> StreamingResponse:
return StreamingResponse(
agent.stream(payload.messages, payload.conversation_id, payload.selected_task_id, payload.provider_id, payload.model_id),
media_type="text/event-stream",
headers={"Cache-Control": "no-cache", "X-Accel-Buffering": "no"},
)
@app.get("/v1/conversations/{conversation_id}")
async def read_conversation(conversation_id: str) -> JSONResponse:
try:
record = store.read_conversation(safe_conversation_id(conversation_id))
except ValueError as error:
raise HTTPException(status_code=400, detail=str(error)) from error
if record is None:
raise HTTPException(status_code=404, detail="Conversation not found")
return JSONResponse(record)
@app.post("/v1/conversations")
async def create_conversation() -> JSONResponse:
return JSONResponse(store.ensure_conversation(None))
@app.patch("/v1/conversations/{conversation_id}")
async def patch_conversation(conversation_id: str, payload: ConversationPatch) -> JSONResponse:
try:
record = store.ensure_conversation(safe_conversation_id(conversation_id), payload.current_task_id, payload.attachments)
except ValueError as error:
raise HTTPException(status_code=400, detail=str(error)) from error
return JSONResponse(record)
@app.post("/v1/uploads")
async def upload_attachment(
file: UploadFile = File(...),
task_id: str | None = Form(default=None),
) -> JSONResponse:
data = await file.read()
filename = file.filename or "attachment"
try:
kind = classify_upload(filename, file.content_type or "", len(data))
task = store.ensure_task(safe_task_id(task_id) if task_id else None, f"Attachment: {filename}")
relative_path, _ = store.write_upload(task["task_id"], filename, data)
extracted_path = ""
if kind == "document":
extracted_path = relative_path + ".txt"
extracted = extract_document_text(data)
store.artifact_path(task["task_id"], extracted_path).write_text(extracted, encoding="utf-8")
record = attachment_record(task["task_id"], filename, file.content_type or "", relative_path, data, kind, extracted_path)
return JSONResponse(record)
except ValueError as error:
raise HTTPException(status_code=400, detail=str(error)) from error
@app.get("/v1/tasks/{task_id}")
async def read_task(task_id: str) -> JSONResponse:
try:
task = store.read_task(safe_task_id(task_id))
except ValueError as error:
raise HTTPException(status_code=400, detail=str(error)) from error
if task is None:
raise HTTPException(status_code=404, detail="Task not found")
return JSONResponse(task)
@app.get("/v1/tasks/{task_id}/artifacts/{artifact_path:path}")
async def read_artifact(task_id: str, artifact_path: str) -> StreamingResponse:
from fastapi.responses import FileResponse
try:
safe_id = safe_task_id(task_id)
path = store.artifact_path(safe_id, artifact_path)
except ValueError as error:
raise HTTPException(status_code=400, detail=str(error)) from error
if not path.is_file():
raise HTTPException(status_code=404, detail="Artifact not found")
return FileResponse(path, filename=path.name)
@app.get("/v1/tasks/{task_id}/parameters")
async def read_parameters(task_id: str) -> JSONResponse:
try:
safe_id = safe_task_id(task_id)
except ValueError as error:
raise HTTPException(status_code=400, detail=str(error)) from error
task = store.read_task(safe_id)
revision_id = str((task or {}).get("current_revision") or "")
revision = next((item for item in (task or {}).get("revisions", []) if item.get("revision_id") == revision_id), None)
relative = str((revision or {}).get("parameters_path") or "")
if not relative:
raise HTTPException(status_code=404, detail="No editable parameters exist for this task")
path = store.artifact_path(safe_id, relative)
if not path.is_file():
raise HTTPException(status_code=404, detail="Parameter contract not found")
return JSONResponse({"task_id": safe_id, "revision_id": revision_id, **json.loads(path.read_text(encoding="utf-8"))})
@app.post("/v1/tasks/{task_id}/parameters")
async def update_parameters(task_id: str, payload: ParameterUpdate) -> JSONResponse:
try:
safe_id = safe_task_id(task_id)
task = store.read_task(safe_id)
current_revision_id = str((task or {}).get("current_revision") or "")
current_path = store.current_cdsl_path(safe_id)
if not task or not current_path or not current_revision_id:
raise ValueError("Task has no successful CDSL revision")
updated, _ = apply_parameter_updates(json.loads(current_path.read_text(encoding="utf-8")), payload.values)
result = build_revision(
settings=settings,
store=store,
task_id=safe_id,
request=f"Parameter update: {', '.join(payload.values)}",
cdsl=updated,
reference_ids=[],
summary="Updated CDSL parameters",
parent_revision_id=current_revision_id,
operation={"type": "parameter_update", "values": payload.values},
)
return JSONResponse(result)
except ValueError as error:
raise HTTPException(status_code=400, detail=str(error)) from error
@app.post("/v1/tasks/{task_id}/modify")
async def modify_task(task_id: str, payload: ModifyRequest) -> JSONResponse:
from app.services.editing import apply_direct_edit
try:
result = apply_direct_edit(settings, store, safe_task_id(task_id), payload.operation, payload.selection, payload.parameters)
return JSONResponse(result)
except ValueError as error:
raise HTTPException(status_code=400, detail=str(error)) from error