- M2-a/b/c: stop/pause/resume/skip dept, plan/rework approval gates, budget & timeout circuit breakers, per-dept cancel and on-demand rework - tier-3: batched event persistence (_EventBuffer), project/dept concurrency semaphores, company_ws event streaming - tier-4: cross-dept deliverable handoff, CEO memory feedback loop, per-dept model override, one-click Markdown report export (RFC5987) - fix: per-department independent SQLAlchemy session in parallel waves (shared Session race across asyncio tasks in company_orchestrator) - fix: rename process_open_fds gauge to app_process_open_fds (collides with prometheus_client ProcessCollector on Linux) - fix: dept role resolution ladder, real token usage & cost tracking, deliverable file capture (mtime scan of dept workspace) - tests: test_company_orchestrator.py 22 cases with mocked LLM - frontend: CompanyControlRoom, companyExecution store, EChart, dept model config dialog, ws proxy, build/typecheck decoupling - docs: virtual company design/usage docs, M2 control plan, deepseek4pro handover docs Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
484 lines
19 KiB
Python
484 lines
19 KiB
Python
"""
|
||
公司项目后台执行器 — 白盒控制室的执行解耦层
|
||
|
||
把公司级流式执行从「请求协程内跑完」解耦为「API 进程内的后台 asyncio 任务」:
|
||
- 执行不再随 SSE/WS 断开而中止(可恢复)
|
||
- 每条编排事件落库 company_project_events(供 WS/HTTP 按 seq 回放)
|
||
- 周期性刷新 project.updated_at(心跳),避免僵尸监管误杀长任务
|
||
|
||
前提:API 为单 uvicorn worker,后台任务与 WS 端点同进程共享 DB/表。
|
||
"""
|
||
from __future__ import annotations
|
||
|
||
import asyncio
|
||
import logging
|
||
import time
|
||
from datetime import datetime
|
||
from typing import Any, Dict, Optional
|
||
|
||
from sqlalchemy import func
|
||
|
||
from app.core.config import settings
|
||
from app.core.database import SessionLocal
|
||
from app.models.company import CompanyProject
|
||
from app.models.company_project_event import CompanyProjectEvent
|
||
from app.services.company_orchestrator import CompanyOrchestrator
|
||
|
||
logger = logging.getLogger(__name__)
|
||
|
||
# 进程内后台任务句柄:防止 asyncio.create_task 结果被 GC;完成后自动清理
|
||
_RUNNING: Dict[str, asyncio.Task] = {}
|
||
|
||
# 三档-b:全局项目并发信号量(lazy,绑定到运行时事件循环)。单 worker 单 loop 单实例。
|
||
_PROJECT_SEM: Optional[asyncio.Semaphore] = None
|
||
|
||
|
||
def _get_project_sem() -> asyncio.Semaphore:
|
||
"""取(缺则建)全局项目并发信号量。lazy 构造避免 import 期绑错事件循环。"""
|
||
global _PROJECT_SEM
|
||
if _PROJECT_SEM is None:
|
||
_PROJECT_SEM = asyncio.Semaphore(max(1, int(settings.MAX_CONCURRENT_COMPANY_PROJECTS)))
|
||
return _PROJECT_SEM
|
||
|
||
|
||
class ProjectControl:
|
||
"""单个运行中项目的进程内控制态(M2-a:暂停/继续/跳过)。
|
||
|
||
与 `_RUNNING` 同为进程内状态:编排任务与 API handler 同进程,靠同一对象通信。
|
||
进程重启则控制态随运行任务一起消失(任务本就死),无持久化需求。
|
||
"""
|
||
|
||
def __init__(self) -> None:
|
||
self.pause_event = asyncio.Event()
|
||
self.pause_event.set() # set = 运行中;clear = 暂停(编排器在检查点 await)
|
||
self.skip_depts: set[str] = set() # 待跳过的(未开始)部门名
|
||
# M2-b 人在环审批门:set = 放行/无门;clear = 编排器在门前 await 等人裁决
|
||
self.approval_event = asyncio.Event()
|
||
self.approval_event.set()
|
||
self.approval_decision: str = "" # 最近一次裁决:"approved" / "rejected"
|
||
self.approval_feedback: str = "" # 人在环注入的反馈文本(返工门用,供下一轮注入部门 prompt)
|
||
self.pending_gate: str = "" # 非空=有门待批,值即门 id:"plan" / "rework"
|
||
# M2-c 部门粒度控制:运行中部门 name→task(供单部门硬取消)
|
||
self.dept_tasks: Dict[str, asyncio.Task] = {}
|
||
self.cancelled_depts: set[str] = set() # 已请求取消的部门名(记录意图)
|
||
|
||
def open_gate(self, gate: str) -> None:
|
||
"""编排器开门等审批:登记门 id、清裁决、阻塞后续 await。"""
|
||
self.pending_gate = gate
|
||
self.approval_decision = ""
|
||
self.approval_feedback = ""
|
||
self.approval_event.clear()
|
||
|
||
def resolve_gate(self, decision: str, feedback: str = "") -> None:
|
||
"""API 端裁决:落决定+可选反馈并放行编排器 await。pending_gate 由编排器过门后清。"""
|
||
self.approval_decision = decision
|
||
self.approval_feedback = feedback or ""
|
||
self.approval_event.set()
|
||
|
||
|
||
# 进程内控制注册表,键为 project_id
|
||
_CONTROL: Dict[str, ProjectControl] = {}
|
||
|
||
|
||
def get_control(project_id: str) -> ProjectControl:
|
||
"""取(缺则建)项目控制态;API 与编排器共用同一对象,幂等。"""
|
||
ctrl = _CONTROL.get(project_id)
|
||
if ctrl is None:
|
||
ctrl = ProjectControl()
|
||
_CONTROL[project_id] = ctrl
|
||
return ctrl
|
||
|
||
|
||
def clear_control(project_id: str) -> None:
|
||
_CONTROL.pop(project_id, None)
|
||
|
||
|
||
def stop(project_id: str) -> bool:
|
||
"""请求停止(硬取消)运行中的项目任务。返回是否有任务被取消。
|
||
|
||
复用 run_company_project 的 CancelledError→stopped 雏形(落 company_terminated)。
|
||
"""
|
||
task = _RUNNING.get(project_id)
|
||
if task is not None and not task.done():
|
||
task.cancel()
|
||
return True
|
||
return False
|
||
|
||
|
||
def cancel_department(project_id: str, name: str) -> bool:
|
||
"""M2-c:硬取消运行中的单个部门任务(不停整个公司)。返回是否有任务被取消。
|
||
|
||
编排器在建部门 task 时把句柄登记进 ctrl.dept_tasks;此处 cancel 触发
|
||
`_run_department_stream` 的 CancelledError 分支,仍发 dept_done 让波次推进、
|
||
下游依赖照常解锁。
|
||
"""
|
||
ctrl = get_control(project_id)
|
||
task = ctrl.dept_tasks.get(name)
|
||
if task is not None and not task.done():
|
||
ctrl.cancelled_depts.add(name)
|
||
task.cancel()
|
||
return True
|
||
return False
|
||
|
||
|
||
def _max_seq(db, project_id: str) -> int:
|
||
"""当前项目已落库事件的最大 seq(无则 0)。返工据此续写,避免前端按 seq 丢弃。"""
|
||
try:
|
||
return db.query(func.max(CompanyProjectEvent.seq)).filter(
|
||
CompanyProjectEvent.project_id == project_id
|
||
).scalar() or 0
|
||
except Exception:
|
||
return 0
|
||
|
||
|
||
_TERMINAL = {"completed", "failed", "stopped"}
|
||
_MAX_STR = 8192 # 单个字符串字段落库上限,防止行爆(如超长 tool_call.input)
|
||
|
||
|
||
def _truncate(obj: Any, limit: int = _MAX_STR):
|
||
"""递归截断过长字符串,控制事件行体积。"""
|
||
if isinstance(obj, str):
|
||
return obj if len(obj) <= limit else obj[:limit] + f"…[truncated {len(obj) - limit} chars]"
|
||
if isinstance(obj, dict):
|
||
return {k: _truncate(v, limit) for k, v in obj.items()}
|
||
if isinstance(obj, list):
|
||
return [_truncate(v, limit) for v in obj]
|
||
return obj
|
||
|
||
|
||
# 三档-a:终态/关键事件类型——立即 force-flush 落地,保证 WS 不截断 finale
|
||
_TERMINAL_EVENT_TYPES = {
|
||
"company_done", "company_terminated", "rework_done",
|
||
"budget_tripped", "timeout_tripped", "error", "review_parse_failed",
|
||
}
|
||
|
||
|
||
class _EventBuffer:
|
||
"""按 run 缓冲事件、批量落库(三档-a)。单事件循环单写者,无需锁。
|
||
|
||
- add() 只入内存;flush() 一次事务 add_all + 单 commit(省远程 MySQL 逐条 RTT)。
|
||
- 心跳折叠进 flush(每 hb_interval 秒顺带 bump updated_at)。
|
||
- 终态/错误事件走 flush_terminal 立即落地,配合 WS 的 2-空轮询宽限防截断 finale。
|
||
- 用 add_all(非 bulk_insert_mappings):模型 id/created_at 是 Python 侧 default,
|
||
add_all 构 ORM 实例默认照常触发;SQLAlchemy 1.4+ 仍渲染多行 INSERT。
|
||
"""
|
||
|
||
def __init__(self, project_id: str) -> None:
|
||
self.project_id = project_id
|
||
self.rows: list[dict] = []
|
||
self.last_flush = time.time()
|
||
self.last_hb = time.time()
|
||
self.max_events = max(1, int(settings.COMPANY_EVENT_FLUSH_MAX_EVENTS))
|
||
self.max_interval = float(settings.COMPANY_EVENT_FLUSH_INTERVAL_SEC)
|
||
self.hb_interval = 5.0
|
||
|
||
def add(self, seq: int, evt: Dict[str, Any]) -> None:
|
||
self.rows.append({
|
||
"project_id": self.project_id,
|
||
"seq": seq,
|
||
"type": str(evt.get("type", "unknown"))[:50],
|
||
"payload": _truncate(evt),
|
||
})
|
||
|
||
def should_flush(self, now: float) -> bool:
|
||
return len(self.rows) >= self.max_events or (
|
||
bool(self.rows) and (now - self.last_flush) >= self.max_interval
|
||
)
|
||
|
||
def flush(self, db, force: bool = False) -> None:
|
||
"""一次事务落盘缓冲事件 + 折叠心跳。失败不阻断执行(沿用旧契约)。"""
|
||
now = time.time()
|
||
hb_due = force or (now - self.last_hb) >= self.hb_interval
|
||
if not self.rows and not hb_due:
|
||
return
|
||
pending = len(self.rows)
|
||
try:
|
||
if self.rows:
|
||
db.add_all([CompanyProjectEvent(**r) for r in self.rows])
|
||
if hb_due:
|
||
db.query(CompanyProject).filter(CompanyProject.id == self.project_id).update(
|
||
{"updated_at": datetime.now()}
|
||
)
|
||
db.commit()
|
||
self.rows.clear()
|
||
self.last_flush = now
|
||
if hb_due:
|
||
self.last_hb = now
|
||
except Exception as e:
|
||
db.rollback()
|
||
logger.warning("事件批量落库失败 [%s, %d 行]: %s", self.project_id, pending, e)
|
||
# 保留 rows 下次重试,但上限 10*N 防持续故障时无限增长(丢最旧)
|
||
cap = self.max_events * 10
|
||
if len(self.rows) > cap:
|
||
dropped = len(self.rows) - cap
|
||
self.rows = self.rows[-cap:]
|
||
logger.warning("事件缓冲超限,丢弃最旧 %d 行 [%s]", dropped, self.project_id)
|
||
|
||
def flush_terminal(self, db, seq: int, evt: Dict[str, Any]) -> None:
|
||
"""终态/错误事件:入队并立即 force-flush(事件 + status-bump 一次 commit 落地)。"""
|
||
self.add(seq, evt)
|
||
self.flush(db, force=True)
|
||
|
||
|
||
def _write_queued(project_id: str, seq: int) -> None:
|
||
"""排队中:临时 session 写一条 queued 事件 + 刷新心跳后关闭(不占用长连接)。"""
|
||
tmp = SessionLocal()
|
||
try:
|
||
_persist_event(tmp, project_id, seq,
|
||
{"type": "queued", "reason": "waiting_slot", "_ts": time.time()})
|
||
_heartbeat(tmp, project_id)
|
||
finally:
|
||
tmp.close()
|
||
|
||
|
||
def _bump_heartbeat(project_id: str) -> None:
|
||
"""排队等待中周期性刷新 updated_at,防僵尸监管误杀(临时 session)。"""
|
||
tmp = SessionLocal()
|
||
try:
|
||
_heartbeat(tmp, project_id)
|
||
finally:
|
||
tmp.close()
|
||
|
||
|
||
def is_running(project_id: str) -> bool:
|
||
task = _RUNNING.get(project_id)
|
||
return task is not None and not task.done()
|
||
|
||
|
||
def spawn(project_id: str, company_id: str, user_id: str, description: str,
|
||
max_rounds: int = 3, ceo_plan: Optional[Dict[str, Any]] = None,
|
||
limits: Optional[Dict[str, Any]] = None) -> asyncio.Task:
|
||
"""在当前事件循环内派发后台执行任务,并登记句柄防 GC。"""
|
||
task = asyncio.create_task(
|
||
run_company_project(project_id, company_id, user_id, description, max_rounds, ceo_plan, limits)
|
||
)
|
||
_RUNNING[project_id] = task
|
||
|
||
def _on_done(t: asyncio.Task) -> None:
|
||
_RUNNING.pop(project_id, None)
|
||
clear_control(project_id) # 清理控制态,避免注册表泄漏
|
||
|
||
task.add_done_callback(_on_done)
|
||
return task
|
||
|
||
|
||
async def run_company_project(
|
||
project_id: str,
|
||
company_id: str,
|
||
user_id: str,
|
||
description: str,
|
||
max_rounds: int = 3,
|
||
ceo_plan: Optional[Dict[str, Any]] = None,
|
||
limits: Optional[Dict[str, Any]] = None,
|
||
) -> None:
|
||
"""后台执行一个公司级项目:驱动 execute_stream,批量事件落库 + 心跳;受全局并发上限约束。"""
|
||
sem = _get_project_sem()
|
||
seq = 0
|
||
acquired = False
|
||
db_orch = None
|
||
db_evt = None
|
||
buf: Optional[_EventBuffer] = None
|
||
try:
|
||
# 三档-b:并发上限 + 真实排队。无空名额时先写 queued 事件、排队期周期心跳。
|
||
if sem.locked():
|
||
seq += 1
|
||
_write_queued(project_id, seq)
|
||
while True:
|
||
try:
|
||
await asyncio.wait_for(sem.acquire(), timeout=settings.COMPANY_QUEUE_HEARTBEAT_SEC)
|
||
acquired = True
|
||
break
|
||
except asyncio.TimeoutError:
|
||
_bump_heartbeat(project_id) # 排队久刷新 updated_at 防僵尸误杀
|
||
|
||
# 名额到手后再开 DB session(排队期零连接,防池死锁)
|
||
db_orch = SessionLocal() # 编排器专用(会被内部并行部门协程共享,沿用既有行为)
|
||
db_evt = SessionLocal() # 事件落库 + 心跳专用,隔离编排器 session 状态
|
||
buf = _EventBuffer(project_id)
|
||
orch = CompanyOrchestrator(db_orch, company_id, user_id)
|
||
async for evt in orch.execute_stream(
|
||
description, max_rounds=max_rounds, ceo_plan=ceo_plan, project_id=project_id,
|
||
limits=limits,
|
||
):
|
||
seq += 1
|
||
buf.add(seq, evt)
|
||
if evt.get("type") in _TERMINAL_EVENT_TYPES:
|
||
buf.flush(db_evt, force=True) # 终态立即落地,防 WS 截断 finale
|
||
else:
|
||
now = time.time()
|
||
if buf.should_flush(now) or (now - buf.last_hb) >= buf.hb_interval:
|
||
buf.flush(db_evt)
|
||
buf.flush(db_evt, force=True) # 收尾兜底:保 company_done 等 finale 落地
|
||
# execute_stream 正常结束时已把 project 状态落 completed/failed(其自身 session)
|
||
logger.info("公司后台执行完成 [%s]: %d 事件", project_id, seq)
|
||
except asyncio.CancelledError:
|
||
if db_evt is None:
|
||
db_evt = SessionLocal()
|
||
if buf is None:
|
||
buf = _EventBuffer(project_id)
|
||
buf.flush(db_evt, force=True)
|
||
_set_status(db_evt, project_id, "stopped")
|
||
buf.flush_terminal(db_evt, seq + 1,
|
||
{"type": "company_terminated", "reason": "cancelled", "_ts": time.time()})
|
||
logger.info("公司后台执行被取消 [%s]", project_id)
|
||
raise
|
||
except Exception as e:
|
||
logger.error("公司后台执行失败 [%s]: %s", project_id, e, exc_info=True)
|
||
if db_evt is None:
|
||
db_evt = SessionLocal()
|
||
if buf is None:
|
||
buf = _EventBuffer(project_id)
|
||
buf.flush(db_evt, force=True)
|
||
_set_status(db_evt, project_id, "failed")
|
||
buf.flush_terminal(db_evt, seq + 1,
|
||
{"type": "error", "message": str(e), "phase": "runner", "_ts": time.time()})
|
||
finally:
|
||
if acquired:
|
||
sem.release() # 先放名额,再关 session
|
||
try:
|
||
if db_orch is not None:
|
||
db_orch.close()
|
||
finally:
|
||
if db_evt is not None:
|
||
db_evt.close()
|
||
|
||
|
||
def spawn_rework(project_id: str, company_id: str, user_id: str, dept_name: str) -> asyncio.Task:
|
||
"""M2-c:派发单部门返工的后台任务(跑完后重跑一个部门)。镜像 spawn 的注册。"""
|
||
task = asyncio.create_task(
|
||
run_department_rework(project_id, company_id, user_id, dept_name)
|
||
)
|
||
_RUNNING[project_id] = task
|
||
|
||
def _on_done(t: asyncio.Task) -> None:
|
||
_RUNNING.pop(project_id, None)
|
||
clear_control(project_id)
|
||
|
||
task.add_done_callback(_on_done)
|
||
return task
|
||
|
||
|
||
async def run_department_rework(
|
||
project_id: str,
|
||
company_id: str,
|
||
user_id: str,
|
||
dept_name: str,
|
||
) -> None:
|
||
"""后台重跑单个部门:seq 续写进同一事件流,替换该部门交付物与打分。
|
||
|
||
结构镜像 run_company_project(并发上限 + 批量落库 + CancelledError/Exception 收尾),
|
||
但不走 execute_stream 全流程,只驱动编排器的 rework_single_department。
|
||
"""
|
||
sem = _get_project_sem()
|
||
acquired = False
|
||
db_orch = None
|
||
db_evt = None
|
||
buf: Optional[_EventBuffer] = None
|
||
seq = 0
|
||
try:
|
||
# 三档-b:返工共用同一全局信号量(=一个部门的完整负载,不该绕过上限)
|
||
if sem.locked():
|
||
tmp = SessionLocal()
|
||
try:
|
||
seq = _max_seq(tmp, project_id) + 1
|
||
_persist_event(tmp, project_id, seq,
|
||
{"type": "queued", "reason": "waiting_slot", "_ts": time.time()})
|
||
_heartbeat(tmp, project_id)
|
||
finally:
|
||
tmp.close()
|
||
while True:
|
||
try:
|
||
await asyncio.wait_for(sem.acquire(), timeout=settings.COMPANY_QUEUE_HEARTBEAT_SEC)
|
||
acquired = True
|
||
break
|
||
except asyncio.TimeoutError:
|
||
_bump_heartbeat(project_id)
|
||
|
||
db_orch = SessionLocal()
|
||
db_evt = SessionLocal()
|
||
buf = _EventBuffer(project_id)
|
||
seq = _max_seq(db_evt, project_id) # 从当前最大 seq(含刚写的 queued)续写
|
||
_set_status(db_evt, project_id, "in_progress")
|
||
seq += 1
|
||
buf.add(seq, {"type": "rework_start", "department_name": dept_name, "_ts": time.time()})
|
||
buf.flush(db_evt, force=True) # 状态刚翻 in_progress,前端要重连,rework_start 尽快可见
|
||
orch = CompanyOrchestrator(db_orch, company_id, user_id)
|
||
orch.auto_approve_files = True
|
||
async for evt in orch.rework_single_department(project_id, dept_name):
|
||
seq += 1
|
||
buf.add(seq, evt)
|
||
if evt.get("type") in _TERMINAL_EVENT_TYPES:
|
||
buf.flush(db_evt, force=True)
|
||
else:
|
||
now = time.time()
|
||
if buf.should_flush(now) or (now - buf.last_hb) >= buf.hb_interval:
|
||
buf.flush(db_evt)
|
||
buf.flush(db_evt, force=True)
|
||
logger.info("单部门返工完成 [%s / %s]", project_id, dept_name)
|
||
except asyncio.CancelledError:
|
||
if db_evt is None:
|
||
db_evt = SessionLocal()
|
||
if buf is None:
|
||
buf = _EventBuffer(project_id)
|
||
buf.flush(db_evt, force=True)
|
||
_set_status(db_evt, project_id, "stopped")
|
||
buf.flush_terminal(db_evt, seq + 1,
|
||
{"type": "company_terminated", "reason": "cancelled", "_ts": time.time()})
|
||
logger.info("单部门返工被取消 [%s / %s]", project_id, dept_name)
|
||
raise
|
||
except Exception as e:
|
||
logger.error("单部门返工失败 [%s / %s]: %s", project_id, dept_name, e, exc_info=True)
|
||
if db_evt is None:
|
||
db_evt = SessionLocal()
|
||
if buf is None:
|
||
buf = _EventBuffer(project_id)
|
||
buf.flush(db_evt, force=True)
|
||
_set_status(db_evt, project_id, "failed")
|
||
buf.flush_terminal(db_evt, seq + 1,
|
||
{"type": "error", "message": str(e), "phase": "rework", "_ts": time.time()})
|
||
finally:
|
||
if acquired:
|
||
sem.release()
|
||
try:
|
||
if db_orch is not None:
|
||
db_orch.close()
|
||
finally:
|
||
if db_evt is not None:
|
||
db_evt.close()
|
||
|
||
|
||
def _persist_event(db, project_id: str, seq: int, evt: Dict[str, Any]) -> None:
|
||
"""落一行事件;失败不阻断执行。"""
|
||
try:
|
||
row = CompanyProjectEvent(
|
||
project_id=project_id,
|
||
seq=seq,
|
||
type=str(evt.get("type", "unknown"))[:50],
|
||
payload=_truncate(evt),
|
||
)
|
||
db.add(row)
|
||
db.commit()
|
||
except Exception as e:
|
||
db.rollback()
|
||
logger.warning("事件落库失败 [%s seq=%s]: %s", project_id, seq, e)
|
||
|
||
|
||
def _heartbeat(db, project_id: str) -> None:
|
||
try:
|
||
db.query(CompanyProject).filter(CompanyProject.id == project_id).update(
|
||
{"updated_at": datetime.now()}
|
||
)
|
||
db.commit()
|
||
except Exception:
|
||
db.rollback()
|
||
|
||
|
||
def _set_status(db, project_id: str, status: str) -> None:
|
||
try:
|
||
db.query(CompanyProject).filter(CompanyProject.id == project_id).update(
|
||
{"status": status, "updated_at": datetime.now()}
|
||
)
|
||
db.commit()
|
||
except Exception:
|
||
db.rollback()
|