181 lines
6.6 KiB
Python
181 lines
6.6 KiB
Python
"""CRUD API for parallel multi-agent task records (并行多智能体任务).
|
||
|
||
A task = one saved run of the standalone parallel multi-agent panel: a task name
|
||
+ brief dispatched to several agents in parallel, plus the resulting message list.
|
||
Bound to the requesting user (not the browser), so it follows the account across
|
||
devices/sessions — same model as roundtable drafts.
|
||
|
||
Routes (prefix ``/api/parallel-agent-tasks``):
|
||
GET "" list current user's tasks (lightweight meta, no messages)
|
||
POST "" create a task
|
||
GET "/{task_id}" one task (full, incl. messages)
|
||
PUT "/{task_id}" partial update (title / brief / agents / messages)
|
||
DELETE "/{task_id}" delete
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import logging
|
||
from datetime import datetime
|
||
from typing import Any
|
||
from uuid import uuid4
|
||
|
||
from fastapi import APIRouter, HTTPException, Request
|
||
from pydantic import BaseModel, Field
|
||
|
||
from deerflow.runtime.user_context import get_effective_user_id
|
||
|
||
logger = logging.getLogger(__name__)
|
||
router = APIRouter(prefix="/api/parallel-agent-tasks", tags=["parallel-agent-tasks"])
|
||
|
||
|
||
# ── schemas ────────────────────────────────────────────────────────────────
|
||
|
||
|
||
class TaskAgentSchema(BaseModel):
|
||
agent_id: str = Field(min_length=1, max_length=128)
|
||
name: str = ""
|
||
avatarType: str | None = None
|
||
|
||
|
||
class TaskResponse(BaseModel):
|
||
id: str
|
||
title: str = ""
|
||
brief: str | None = None
|
||
agents: list[TaskAgentSchema] = []
|
||
# messages is an opaque dialogue list (mirrors the frontend Dialogue[] shape);
|
||
# absent on list responses (lightweight), present on get.
|
||
messages: list[dict[str, Any]] | None = None
|
||
created_at: datetime | str | None = None
|
||
updated_at: datetime | str | None = None
|
||
|
||
|
||
class TaskListResponse(BaseModel):
|
||
tasks: list[TaskResponse]
|
||
|
||
|
||
class TaskCreateRequest(BaseModel):
|
||
id: str | None = Field(default=None, max_length=64)
|
||
title: str = Field(default="", max_length=512)
|
||
brief: str | None = None
|
||
agents: list[TaskAgentSchema] = []
|
||
messages: list[dict[str, Any]] = []
|
||
|
||
|
||
class TaskUpdateRequest(BaseModel):
|
||
# All optional → callers send only what changed.
|
||
title: str | None = Field(default=None, max_length=512)
|
||
brief: str | None = None
|
||
agents: list[TaskAgentSchema] | None = None
|
||
messages: list[dict[str, Any]] | None = None
|
||
|
||
|
||
class TaskEnsureRequest(BaseModel):
|
||
"""get-or-create:宿主传的 taskId 不固定,新 id 直接**初始化**一条空记录返回(不 404)。"""
|
||
id: str = Field(min_length=1, max_length=64)
|
||
title: str = Field(default="", max_length=512)
|
||
brief: str | None = None
|
||
agents: list[TaskAgentSchema] = []
|
||
|
||
|
||
# ── helpers ──────────────────────────────────────────────────────────────
|
||
|
||
|
||
def _current_user_id(request: Request) -> str:
|
||
user = getattr(request.state, "user", None)
|
||
if user is not None:
|
||
return str(user.id)
|
||
return get_effective_user_id()
|
||
|
||
|
||
def _get_store(request: Request):
|
||
store = getattr(request.app.state, "parallel_agent_task_store", None)
|
||
if store is None:
|
||
raise HTTPException(status_code=503, detail="Parallel agent task store not available")
|
||
return store
|
||
|
||
|
||
# ── routes ───────────────────────────────────────────────────────────────
|
||
|
||
|
||
@router.get("", response_model=TaskListResponse)
|
||
async def list_tasks(request: Request) -> TaskListResponse:
|
||
store = _get_store(request)
|
||
user_id = _current_user_id(request)
|
||
rows = await store.list_tasks(user_id)
|
||
return TaskListResponse(tasks=[TaskResponse(**r) for r in rows])
|
||
|
||
|
||
@router.post("", response_model=TaskResponse, status_code=201)
|
||
async def create_task(request: Request, body: TaskCreateRequest) -> TaskResponse:
|
||
store = _get_store(request)
|
||
user_id = _current_user_id(request)
|
||
row = await store.create_task(
|
||
user_id,
|
||
{
|
||
"id": (body.id or uuid4().hex),
|
||
"title": body.title.strip(),
|
||
"brief": body.brief,
|
||
"agents": [a.model_dump() for a in body.agents],
|
||
"messages": body.messages,
|
||
},
|
||
)
|
||
return TaskResponse(**row)
|
||
|
||
|
||
@router.post("/ensure", response_model=TaskResponse)
|
||
async def ensure_task(request: Request, body: TaskEnsureRequest) -> TaskResponse:
|
||
"""存在则返回(含 messages 供回显);不存在则用该 id **初始化**一条空记录返回。
|
||
|
||
宿主传的 taskId 不固定:新 id 不再 404,而是即时建好、可直接回显/保存。幂等。
|
||
"""
|
||
store = _get_store(request)
|
||
user_id = _current_user_id(request)
|
||
existing = await store.get_task(body.id, user_id)
|
||
if existing is not None:
|
||
return TaskResponse(**existing)
|
||
row = await store.create_task(
|
||
user_id,
|
||
{
|
||
"id": body.id,
|
||
"title": body.title.strip(),
|
||
"brief": body.brief,
|
||
"agents": [a.model_dump() for a in body.agents],
|
||
"messages": [],
|
||
},
|
||
)
|
||
return TaskResponse(**row)
|
||
|
||
|
||
@router.get("/{task_id}", response_model=TaskResponse)
|
||
async def get_task(request: Request, task_id: str) -> TaskResponse:
|
||
store = _get_store(request)
|
||
user_id = _current_user_id(request)
|
||
row = await store.get_task(task_id, user_id)
|
||
if row is None:
|
||
raise HTTPException(status_code=404, detail="Task not found")
|
||
return TaskResponse(**row)
|
||
|
||
|
||
@router.put("/{task_id}", response_model=TaskResponse)
|
||
async def update_task(request: Request, task_id: str, body: TaskUpdateRequest) -> TaskResponse:
|
||
store = _get_store(request)
|
||
user_id = _current_user_id(request)
|
||
fields = body.model_dump(exclude_unset=True)
|
||
# Normalize agents (pydantic models) to plain dicts for JSON storage.
|
||
if "agents" in fields and fields["agents"] is not None:
|
||
fields["agents"] = [a if isinstance(a, dict) else a.model_dump() for a in body.agents or []]
|
||
row = await store.update_task(task_id, user_id, **fields)
|
||
if row is None:
|
||
raise HTTPException(status_code=404, detail="Task not found")
|
||
return TaskResponse(**row)
|
||
|
||
|
||
@router.delete("/{task_id}", status_code=204)
|
||
async def delete_task(request: Request, task_id: str) -> None:
|
||
store = _get_store(request)
|
||
user_id = _current_user_id(request)
|
||
deleted = await store.delete_task(task_id, user_id)
|
||
if not deleted:
|
||
raise HTTPException(status_code=404, detail="Task not found")
|