deerflow-code/offline-backend-20260512/backend/app/report_collaboration/planning/service.py
2026-09-07 18:24:55 +08:00

153 lines
6.9 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""Fulfill a durable plan_request: catalog → select → assemble → validate → persist.
Does not start a run. Rule-based until the BE-7 model adapter is wired.
"""
from __future__ import annotations
from typing import Any
from uuid import uuid4
from app.report_collaboration.contracts.requirements import ReportRequirementSnapshot
from app.report_collaboration.planning.assembler import PlanAssembler
from app.report_collaboration.planning.catalog import project_agent_catalog
from app.report_collaboration.planning.controller import PlanController
from app.report_collaboration.planning.validator import apply_validation, validate_assembled_plan
from app.report_collaboration.requirement.defaults import COORDINATOR_NAME, DEFAULT_TITLE
from deerflow.config.app_config import get_app_config
from deerflow.config.report_collaboration_config import ReportCollaborationConfig
from deerflow.persistence.report_collaboration import (
ReportCollaborationConflictError,
ReportCollaborationStore,
ReportCollaborationValidationError,
)
_BUSY = frozenset({"running", "awaiting_input", "reviewing", "completed", "failed", "cancelled"})
def requirement_from_session(record: dict[str, Any]) -> ReportRequirementSnapshot:
raw = record.get("requirement_json")
snapshot: dict[str, Any] | None = None
if isinstance(raw, dict) and isinstance(raw.get("snapshot"), dict):
snapshot = dict(raw["snapshot"])
elif isinstance(raw, dict) and raw.get("topic"):
snapshot = dict(raw)
if snapshot is None or not str(snapshot.get("topic") or "").strip() or snapshot.get("topic") == "未指定主题":
title = str(record.get("title") or "").strip()
if title and title not in {DEFAULT_TITLE, ""}:
snapshot = {"topic": title, "required_angles": [], "excluded_angles": [], "revision": max(1, int(record.get("requirement_revision") or 1))}
else:
raise ReportCollaborationValidationError("REQUIREMENT_INCOMPLETE", "请先说明报告主题,再生成协作方案")
snapshot.setdefault("required_angles", [])
snapshot.setdefault("excluded_angles", [])
snapshot["revision"] = max(1, int(snapshot.get("revision") or record.get("requirement_revision") or 1))
return ReportRequirementSnapshot.model_validate(snapshot)
class PlanProposalService:
def __init__(
self,
store: ReportCollaborationStore,
*,
agent_store: Any | None = None,
config: ReportCollaborationConfig | None = None,
) -> None:
self._store = store
self._agent_store = agent_store
self._config = config
self._controller = PlanController()
self._assembler = PlanAssembler()
def _cfg(self) -> ReportCollaborationConfig:
return self._config or get_app_config().report_collaboration
async def fulfill(self, session_id: str, *, user_id: str, command_id: str, idempotency_key: str) -> list[dict[str, Any]]:
record = await self._store.get_session_record(session_id)
if record is None:
raise ReportCollaborationValidationError("NOT_FOUND", "会话不存在")
if record.get("selected_plan_id"):
raise ReportCollaborationConflictError("PLAN_ALREADY_SELECTED", "会话已选定方案,无法重新规划", current_revision=int(record.get("requirement_revision") or 0))
if record.get("status") in _BUSY:
raise ReportCollaborationConflictError("SESSION_BUSY", "当前会话状态不允许生成方案", current_revision=int(record.get("requirement_revision") or 0))
try:
return await self._build(session_id, user_id=user_id, command_id=command_id, idempotency_key=idempotency_key, record=record)
except Exception as exc:
await self._fail_command(session_id, command_id, exc)
raise
async def _fail_command(self, session_id: str, command_id: str, exc: Exception) -> None:
try:
await self._store.update_command_status(command_id, status="failed", error=str(exc)[:500])
record = await self._store.get_session_record(session_id)
if record is None or record.get("status") != "planning":
return
plans = await self._store.list_plans(session_id)
if any(item.get("status") in {"proposed", "selected"} for item in plans):
restore = "proposal_ready"
elif int(record.get("requirement_revision") or 0) > 0:
restore = "clarifying"
else:
restore = "empty"
await self._store.save_requirement_state(
session_id,
requirement_json=record.get("requirement_json") or {},
requirement_revision=int(record.get("requirement_revision") or 0),
status=restore,
)
except Exception:
return
async def _build(
self,
session_id: str,
*,
user_id: str,
command_id: str,
idempotency_key: str,
record: dict[str, Any],
) -> list[dict[str, Any]]:
requirement = requirement_from_session(record)
cfg = self._cfg()
catalog = await project_agent_catalog(user_id=user_id, config=cfg, agent_store=self._agent_store)
strategies = self._controller.propose(
requirement,
max_parallel_tasks=cfg.max_parallel_tasks,
max_team_members=cfg.max_team_members,
)
group_id = f"rpg_{uuid4().hex[:12]}"
assembled_list = [
self._assembler.assemble(
strategy,
requirement=requirement,
catalog=catalog,
config=cfg,
proposal_group_id=group_id,
revision=1,
)
for strategy in strategies
]
persisted: list[dict[str, Any]] = []
for assembled in assembled_list:
validation = validate_assembled_plan(assembled, catalog=catalog, config=cfg)
apply_validation(assembled, validation)
if not validation.ok:
raise ReportCollaborationValidationError("PLAN_INVALID", ";".join(validation.errors))
persisted.append(assembled.candidate)
await self._store.supersede_proposed_plans(session_id)
saved: list[dict[str, Any]] = []
for candidate in persisted:
saved.append(await self._store.insert_plan(session_id, candidate.model_dump()))
await self._store.update_command_status(command_id, status="completed")
await self._store.create_message(
session_id,
role="ai",
content="已根据当前需求生成 3 个候选协作方案。请选择其中一个后再开始执行;选择方案不会自动开跑。",
idempotency_key=f"{idempotency_key}:plans",
name=COORDINATOR_NAME,
metadata={"plan_request_command_id": command_id, "proposal_group_id": group_id},
completed=True,
)
return saved