deerflow-code/offline-backend-20260512/backend/app/gateway/routers/task_reports.py
2026-09-07 18:24:55 +08:00

267 lines
10 KiB
Python

"""Task-scoped situation-report detail API.
The report payload intentionally mirrors the existing ``selectSqReport``
Consumer contract:
``{state, msg, data: [{id, taskId, createTime, updateTime, sessionId,
contentJson, categoryType}]}``. ``contentJson`` stays a JSON string so the
dashboard parser can continue to deserialize it category by category.
Routes (all authenticated by the Gateway middleware):
GET /api/task-reports/by-task/{task_id}
POST /api/task-reports/by-task/{task_id}
PUT /api/task-reports/by-task/{task_id}
"""
from __future__ import annotations
import json
from datetime import UTC, datetime
from typing import Any
from uuid import uuid4
from fastapi import APIRouter, Body, HTTPException, Request
from pydantic import AliasChoices, BaseModel, ConfigDict, Field, StrictStr
from app.gateway.deps import get_current_user
from deerflow.persistence.task_reports import TaskReportStore
router = APIRouter(prefix="/api/task-reports", tags=["task-reports"])
class TaskReportRecordInput(BaseModel):
"""One ``data`` entry in the legacy-compatible report envelope."""
model_config = ConfigDict(populate_by_name=True)
id: StrictStr | None = Field(default=None, max_length=128)
task_id: str | int | None = Field(
default=None,
validation_alias=AliasChoices("taskId", "task_id"),
)
create_time: StrictStr | None = Field(
default=None,
validation_alias=AliasChoices("createTime", "create_time"),
)
update_time: StrictStr | None = Field(
default=None,
validation_alias=AliasChoices("updateTime", "update_time"),
)
session_id: StrictStr = Field(
default="",
max_length=255,
validation_alias=AliasChoices("sessionId", "session_id"),
)
content_json: StrictStr = Field(
...,
validation_alias=AliasChoices("contentJson", "content_json"),
)
category_type: StrictStr = Field(
...,
max_length=255,
validation_alias=AliasChoices("categoryType", "category_type"),
)
class TaskReportWriteRequest(BaseModel):
data: list[TaskReportRecordInput] = Field(min_length=1)
class TaskReportRecordResponse(BaseModel):
model_config = ConfigDict(populate_by_name=True)
id: str
task_id: int | str = Field(serialization_alias="taskId")
create_time: str = Field(serialization_alias="createTime")
update_time: str = Field(serialization_alias="updateTime")
session_id: str = Field(serialization_alias="sessionId")
content_json: str = Field(serialization_alias="contentJson")
category_type: str = Field(serialization_alias="categoryType")
class TaskReportResponse(BaseModel):
state: str = "200"
msg: str = "操作成功!"
data: list[TaskReportRecordResponse]
def _get_store(request: Request) -> TaskReportStore:
store = getattr(request.app.state, "task_report_store", None)
if store is None:
raise HTTPException(status_code=503, detail="Task report store not available")
return store
async def _mark_task_analysis_completed(request: Request, task_id: str) -> None:
"""Reflect a non-empty saved report in the TaskCOP task-list status."""
task_store = getattr(request.app.state, "taskcop_task_store", None)
if task_store is not None:
await task_store.mark_analysis_completed(task_id)
def _public_task_id(task_id: str) -> int | str:
"""Keep numeric task ids numeric, matching the supplied mock payload."""
return int(task_id) if task_id.isdecimal() and task_id == str(int(task_id)) else task_id
def _to_response(records: list[dict[str, Any]]) -> TaskReportResponse:
return TaskReportResponse(
data=[
TaskReportRecordResponse(
id=str(record["id"]),
task_id=record["taskId"],
create_time=str(record["createTime"]),
update_time=str(record["updateTime"]),
session_id=str(record["sessionId"]),
content_json=str(record["contentJson"]),
category_type=str(record["categoryType"]),
)
for record in records
]
)
def _normalize_records(
task_id: str,
data: list[TaskReportRecordInput],
previous: list[dict[str, Any]] | None = None,
) -> list[dict[str, Any]]:
"""Return the exact public record shape and preserve original create times."""
existing_by_id = {str(item.get("id")): item for item in previous or [] if item.get("id")}
now = datetime.now(UTC).isoformat()
normalized: list[dict[str, Any]] = []
seen_ids: set[str] = set()
for item in data:
if item.task_id is not None and str(item.task_id) != task_id:
raise HTTPException(status_code=422, detail="data.taskId must match the path task id")
record_id = item.id or uuid4().hex
if record_id in seen_ids:
raise HTTPException(status_code=422, detail="data contains duplicate id values")
seen_ids.add(record_id)
before = existing_by_id.get(record_id, {})
normalized.append(
{
"id": record_id,
"taskId": _public_task_id(task_id),
# Importing the supplied mock retains its createTime. On later
# edits, matching ids retain their original createTime.
"createTime": before.get("createTime") or item.create_time or now,
"updateTime": now,
"sessionId": item.session_id,
"contentJson": item.content_json,
"categoryType": item.category_type,
}
)
return normalized
def _import_records(payload: Any, task_id: str) -> list[TaskReportRecordInput]:
"""Read the supplied mock envelope and bind every record to ``task_id``.
Detail files are frequently exported from another task. The import page
deliberately makes the selected task authoritative, so stale taskId values
inside the file cannot write or reject the intended target.
"""
if isinstance(payload, list):
raw_records = payload
elif isinstance(payload, dict):
raw_records = payload.get("data", payload.get("records"))
else:
raw_records = None
if not isinstance(raw_records, list) or not raw_records:
raise HTTPException(status_code=422, detail="详情导入文件必须包含非空的 data 或 records 数组")
records: list[TaskReportRecordInput] = []
for index, raw in enumerate(raw_records, start=1):
if not isinstance(raw, dict):
raise HTTPException(status_code=422, detail=f"第 {index} 条详情数据不是对象")
value = dict(raw)
content_key = "contentJson" if "contentJson" in value else "content_json"
if content_key in value and not isinstance(value[content_key], str):
value[content_key] = json.dumps(value[content_key], ensure_ascii=False, separators=(",", ":"))
# Prefer the public key so AliasChoices resolves this value consistently.
value.pop("task_id", None)
value["taskId"] = task_id
try:
records.append(TaskReportRecordInput.model_validate(value))
except Exception as error:
raise HTTPException(status_code=422, detail=f"第 {index} 条详情数据格式不正确:{error}") from error
return records
@router.get("/by-task/{task_id}", response_model=TaskReportResponse)
async def get_task_report(request: Request, task_id: str) -> TaskReportResponse:
"""Read a task report. Missing reports use the same envelope with ``data: []``."""
row = await _get_store(request).get_by_task(task_id)
return _to_response(row.get("records", []) if row is not None else [])
@router.post("/by-task/{task_id}", response_model=TaskReportResponse)
async def create_task_report(
request: Request,
task_id: str,
body: TaskReportWriteRequest,
) -> TaskReportResponse:
"""Save a task's first report document; reject duplicate initial saves."""
records = _normalize_records(task_id, body.data)
row = await _get_store(request).create_by_task(
task_id,
records,
updated_by=await get_current_user(request),
)
if row is None:
raise HTTPException(status_code=409, detail="Task report already exists; use PUT to modify it")
await _mark_task_analysis_completed(request, task_id)
return _to_response(row["records"])
@router.put("/by-task/{task_id}", response_model=TaskReportResponse)
async def update_task_report(
request: Request,
task_id: str,
body: TaskReportWriteRequest,
) -> TaskReportResponse:
"""Modify a report by atomically replacing its complete ``data`` array."""
store = _get_store(request)
current = await store.get_by_task(task_id)
if current is None:
raise HTTPException(status_code=404, detail="Task report not found; use POST to save it first")
records = _normalize_records(task_id, body.data, current.get("records", []))
row = await store.update_by_task(
task_id,
records,
updated_by=await get_current_user(request),
)
if row is None: # defensive: another process could remove it between read and write
raise HTTPException(status_code=404, detail="Task report not found")
await _mark_task_analysis_completed(request, task_id)
return _to_response(row["records"])
@router.post("/import/by-task/{task_id}", response_model=TaskReportResponse)
async def import_task_report(
request: Request, task_id: str, payload: Any = Body(...)
) -> TaskReportResponse:
"""Import and replace a selected task's complete report document."""
task_id = task_id.strip()
if not task_id:
raise HTTPException(status_code=422, detail="task id is required")
task_store = getattr(request.app.state, "taskcop_task_store", None)
if task_store is not None and await task_store.get_task(task_id) is None:
raise HTTPException(status_code=404, detail="任务不存在,请先导入或创建任务")
store = _get_store(request)
current = await store.get_by_task(task_id)
records = _normalize_records(task_id, _import_records(payload, task_id), current.get("records", []) if current else None)
if current is None:
row = await store.create_by_task(task_id, records, updated_by=await get_current_user(request))
else:
row = await store.update_by_task(task_id, records, updated_by=await get_current_user(request))
if row is None: # defensive: a concurrent delete/import may have won the race
raise HTTPException(status_code=409, detail="详情导入时发生并发冲突,请重试")
await _mark_task_analysis_completed(request, task_id)
return _to_response(row["records"])