Files
aiagent/backend/app/services/company_runner.py
renjianbo 495eafcd78 feat: virtual company M2 control + tier-3/4 capabilities + per-dept session fix
- 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>
2026-07-26 18:16:39 +08:00

484 lines
19 KiB
Python
Raw 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.
"""
公司项目后台执行器 — 白盒控制室的执行解耦层
把公司级流式执行从「请求协程内跑完」解耦为「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()