"""SQL repository for position-collaboration sessions and chain nodes.""" from __future__ import annotations import hashlib import json from datetime import UTC, datetime from typing import Any from uuid import NAMESPACE_URL, uuid4, uuid5 from sqlalchemy import delete, select from sqlalchemy.exc import IntegrityError from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker from deerflow.persistence.position_roundtable.delivery import ( is_intent_position_node, is_markdown_artifact, markdown_delivery_changed, session_intent_position_id, ) from deerflow.persistence.position_roundtable.model import ( PositionRoundtableCommandRow, PositionRoundtableNodeRow, PositionRoundtableSessionRow, ) _UNSET = object() def _dumps(value: Any) -> str | None: if value is None: return None return json.dumps(value, ensure_ascii=False) def _loads(value: Any) -> Any: if not isinstance(value, str) or not value: return None try: return json.loads(value) except Exception: return None def _time(value: Any) -> str | Any: return value.isoformat() if isinstance(value, datetime) else value def _has_material(row: PositionRoundtableNodeRow) -> bool: return row.revision > 0 or bool(row.latest_answer) or bool(_loads(row.artifact_manifest)) def _has_confirmed_intent_snapshot(intent: Any) -> bool: """Return whether a persisted Step-1 snapshot may activate a chain. This duplicates only the storage-level shape check used by the router so activation can re-check the prerequisite *after* locking the session row. A router-only pre-check leaves a TOCTOU window where another user can overwrite the intent before the chain topology is frozen. """ if not isinstance(intent, dict): return False required_sections = ("coreGoals", "riskWarnings", "keyPoints", "strategicSignificance") if any(key in intent for key in required_sections): if not all(isinstance(intent.get(key), list) for key in required_sections): return False return bool(intent["coreGoals"] or intent["strategicSignificance"]) objective = intent.get("objective") or intent.get("goal") return bool(isinstance(objective, str) and objective.strip() and isinstance(intent.get("constraints"), list)) class NodeConcurrentWriteError(Exception): """节点被并发改写,调用方持有的 version 已过期(乐观锁冲突)。 由 ``update_node(expected_version=...)`` / ``recompute_chain_state`` 在检测到 version 不匹配时抛出。调用方可捕获后重新读取最新状态重算(幂等重入)。 """ def __init__( self, node_key: str, *, expected: int | None = None, actual: int | None = None, ) -> None: self.node_key = node_key self.expected = expected self.actual = actual super().__init__( f"concurrent write detected on position node {node_key}" + (f" (expected version {expected}, got {actual})" if expected is not None else "") ) class SessionWriteConflictError(Exception): """Session-level state changed after a route's optimistic pre-check. The public routes translate this into HTTP 409. Keeping the exception at the persistence boundary prevents a stale intent/task save from reviving an archived session or changing the frozen input of an active chain. """ def __init__(self, detail: str) -> None: self.detail = detail super().__init__(detail) class CommandConflictError(Exception): """命令在锁内判定为业务冲突(→ HTTP 409),事务已回滚、未改任何状态。 与「幂等重放」严格区分:重放返回上次成功结果;冲突说明状态机拒绝该命令 (节点未解锁 / 已绑定其它线程 / 同 command_id 携带了不同参数等)。 ``detail`` 是面向用户的提示,router 直接转成 409。 """ def __init__(self, detail: str) -> None: self.detail = detail super().__init__(detail) def _payload_hash(command_type: str, payload: dict[str, Any]) -> str: """规范化命令参数的 sha256,用于「同 command_id 是否同一参数」判定。""" normalized = json.dumps( {"type": command_type, "payload": payload}, ensure_ascii=False, sort_keys=True, separators=(",", ":"), ) return hashlib.sha256(normalized.encode("utf-8")).hexdigest() class PositionRoundtableRepository: def __init__(self, session_factory: async_sessionmaker[AsyncSession]) -> None: self._sf = session_factory @staticmethod def _can_access(row: PositionRoundtableSessionRow, user_id: str) -> bool: """History is partitioned by creator; even task-bound rows stay private.""" return row.user_id == user_id @staticmethod def _session_dict(row: PositionRoundtableSessionRow) -> dict[str, Any]: return { "id": row.id, "external_task_id": row.external_task_id, "task_scoped": bool(row.task_scoped), "task_snapshot": _loads(row.task_snapshot) or {}, "chain_id": row.chain_id, "chain_snapshot": _loads(row.chain_snapshot), "intent_thread_id": row.intent_thread_id, "intent_snapshot": _loads(row.intent_snapshot), "intent_conversation": _loads(row.intent_conversation), "summary_thread_id": row.summary_thread_id, "summary_snapshot": _loads(row.summary_snapshot), "summary_conversation": _loads(row.summary_conversation), "action_plan_thread_id": row.action_plan_thread_id, "action_plan_snapshot": _loads(row.action_plan_snapshot), "action_plan_conversation": _loads(row.action_plan_conversation), "status": row.status, "version": row.version, "created_at": _time(row.created_at), "updated_at": _time(row.updated_at), } @staticmethod def _node_dict(row: PositionRoundtableNodeRow) -> dict[str, Any]: return { "id": row.id, "session_id": row.session_id, "node_key": row.node_key, "stage_index": row.stage_index, "seat_index": row.seat_index, "agent_id": row.agent_id, "position_id": row.position_id, "thread_id": row.thread_id, "status": row.status, "latest_answer": row.latest_answer, "artifact_manifest": _loads(row.artifact_manifest) or [], "conversation_snapshot": _loads(row.conversation_snapshot), "revision": row.revision, "version": row.version, "upstream_revision": _loads(row.upstream_revision), "rejection_count": row.rejection_count or 0, "last_rejection": _loads(row.last_rejection), "invalidated_by": row.invalidated_by, "created_at": _time(row.created_at), "updated_at": _time(row.updated_at), } async def create_session( self, user_id: str, *, external_task_id: str | None, task_snapshot: dict[str, Any] | None, task_scoped: bool | None = None, ) -> dict[str, Any]: now = datetime.now(UTC) row = PositionRoundtableSessionRow( id=uuid4().hex, user_id=user_id, external_task_id=external_task_id or None, task_scoped=bool(external_task_id) if task_scoped is None else task_scoped, task_snapshot=_dumps(task_snapshot or {}), status="intent_pending", created_at=now, updated_at=now, ) async with self._sf() as session: session.add(row) await session.commit() await session.refresh(row) return self._session_dict(row) async def ensure_demo_session( self, user_id: str, *, demo_key: str, session_payload: dict[str, Any], nodes: list[dict[str, Any]], ) -> tuple[dict[str, Any], list[dict[str, Any]]]: """Create one deterministic, user-scoped demonstration session. The browser owns the canonical JSON fixture so the same payload can be rendered locally when the gateway is unavailable. When the gateway is reachable this method imports that payload exactly once. Stable ids make repeated initialisation and React StrictMode double effects safe, while the user id in every UUID seed avoids cross-account collisions on the globally unique LangGraph thread columns. """ def stable_id(label: str) -> str: return uuid5( NAMESPACE_URL, f"deerflow:position-roundtable-demo:{user_id}:{demo_key}:{label}", ).hex session_id = stable_id("session") existing = await self.get_session(session_id, user_id) if existing is not None: return existing, (await self.list_nodes(session_id, user_id)) or [] now = datetime.now(UTC) intent_thread_id = ( stable_id("intent-thread") if session_payload.get("intent_thread_id") else None ) summary_thread_id = ( stable_id("summary-thread") if session_payload.get("summary_thread_id") else None ) action_plan_thread_id = ( stable_id("action-plan-thread") if session_payload.get("action_plan_thread_id") else None ) def with_thread_id(value: Any, thread_id: str | None) -> Any: if not isinstance(value, dict) or not thread_id: return value return {**value, "thread_id": thread_id} row = PositionRoundtableSessionRow( id=session_id, user_id=user_id, external_task_id=str(session_payload.get("external_task_id") or "").strip() or None, task_scoped=False, task_snapshot=_dumps(session_payload.get("task_snapshot") or {}), chain_id=session_payload.get("chain_id") or None, chain_snapshot=_dumps(session_payload.get("chain_snapshot")), intent_thread_id=intent_thread_id, intent_snapshot=_dumps(session_payload.get("intent_snapshot")), intent_conversation=_dumps(session_payload.get("intent_conversation")), summary_thread_id=summary_thread_id, summary_snapshot=_dumps( with_thread_id(session_payload.get("summary_snapshot"), summary_thread_id) ), summary_conversation=_dumps(session_payload.get("summary_conversation")), action_plan_thread_id=action_plan_thread_id, action_plan_snapshot=_dumps( with_thread_id( session_payload.get("action_plan_snapshot"), action_plan_thread_id ) ), action_plan_conversation=_dumps( session_payload.get("action_plan_conversation") ), status=str(session_payload.get("status") or "completed"), created_at=now, updated_at=now, ) node_rows: list[PositionRoundtableNodeRow] = [] for index, node in enumerate(nodes): node_key = str(node.get("node_key") or f"demo-node-{index + 1}") node_rows.append( PositionRoundtableNodeRow( id=stable_id(f"node:{node_key}"), session_id=session_id, node_key=node_key, stage_index=int(node.get("stage_index") or 0), seat_index=int(node.get("seat_index") or 0), agent_id=str(node.get("agent_id") or f"demo-agent-{index + 1}"), position_id=node.get("position_id") or None, thread_id=( stable_id(f"node-thread:{node_key}") if node.get("thread_id") else None ), status=str(node.get("status") or "done"), latest_answer=node.get("latest_answer") or None, artifact_manifest=_dumps(node.get("artifact_manifest") or []), conversation_snapshot=_dumps(node.get("conversation_snapshot")), revision=int(node.get("revision") or 0), upstream_revision=_dumps(node.get("upstream_revision")), rejection_count=int(node.get("rejection_count") or 0), last_rejection=_dumps(node.get("last_rejection")), invalidated_by=node.get("invalidated_by") or None, created_at=now, updated_at=now, ) ) try: async with self._sf() as session: session.add(row) session.add_all(node_rows) await session.commit() await session.refresh(row) for node_row in node_rows: await session.refresh(node_row) except IntegrityError: # Another tab may have seeded the same user/demo key between the # initial lookup and commit. The deterministic row is the desired # result, so recover it instead of returning a spurious 500. existing = await self.get_session(session_id, user_id) if existing is None: raise return existing, (await self.list_nodes(session_id, user_id)) or [] return self._session_dict(row), [ self._node_dict(node_row) for node_row in node_rows ] async def list_sessions( self, user_id: str, *, external_task_id: str | None = None, limit: int = 50 ) -> list[dict[str, Any]]: """List the caller's own task history or legacy personal history.""" if external_task_id: stmt = select(PositionRoundtableSessionRow).where( PositionRoundtableSessionRow.external_task_id == external_task_id, PositionRoundtableSessionRow.task_scoped.is_(True), PositionRoundtableSessionRow.user_id == user_id, ) else: stmt = select(PositionRoundtableSessionRow).where( PositionRoundtableSessionRow.user_id == user_id, PositionRoundtableSessionRow.task_scoped.is_(False), ) stmt = stmt.order_by(PositionRoundtableSessionRow.updated_at.desc()).limit(limit) async with self._sf() as session: rows = (await session.execute(stmt)).scalars().all() return [self._session_dict(row) for row in rows] async def get_session(self, session_id: str, user_id: str) -> dict[str, Any] | None: async with self._sf() as session: row = await session.get(PositionRoundtableSessionRow, session_id) if row is None or not self._can_access(row, user_id): return None return self._session_dict(row) async def update_session( self, session_id: str, user_id: str, *, task_snapshot: dict[str, Any] | object = _UNSET, intent_thread_id: str | None | object = _UNSET, intent_snapshot: dict[str, Any] | None | object = _UNSET, intent_conversation: dict[str, Any] | None | object = _UNSET, summary_thread_id: str | None | object = _UNSET, summary_snapshot: dict[str, Any] | None | object = _UNSET, summary_conversation: dict[str, Any] | None | object = _UNSET, action_plan_thread_id: str | None | object = _UNSET, action_plan_snapshot: dict[str, Any] | None | object = _UNSET, action_plan_conversation: dict[str, Any] | None | object = _UNSET, status: str | object = _UNSET, allowed_statuses: set[str] | None = None, ) -> dict[str, Any] | None: """Update one session snapshot under the same row lock as commands. Snapshot persistence used to do an unlocked ORM ``get`` followed by a write. That permitted a stale "confirm intent" save to land after a concurrent activation/archival, or allowed two different snapshot surfaces to overwrite each other's unseen values. Acquiring the parent row lock serializes all session writes with node/session commands. ``allowed_statuses`` additionally freezes task/intent inputs once the business chain has become active. """ async with self._sf() as session: row = ( await session.execute( select(PositionRoundtableSessionRow) .where(PositionRoundtableSessionRow.id == session_id) .with_for_update() ) ).scalar_one_or_none() if row is None or not self._can_access(row, user_id): return None if row.status == "archived": raise SessionWriteConflictError("该岗位协同会话已归档,仅可查看历史记录") if allowed_statuses is not None and row.status not in allowed_statuses: raise SessionWriteConflictError("会话已启动,不能再修改已冻结的任务或意图") if task_snapshot is not _UNSET: row.task_snapshot = _dumps(task_snapshot) if intent_thread_id is not _UNSET: row.intent_thread_id = intent_thread_id or None if intent_snapshot is not _UNSET: row.intent_snapshot = _dumps(intent_snapshot) if intent_conversation is not _UNSET: row.intent_conversation = _dumps(intent_conversation) if summary_thread_id is not _UNSET: row.summary_thread_id = summary_thread_id or None if summary_snapshot is not _UNSET: row.summary_snapshot = _dumps(summary_snapshot) if summary_conversation is not _UNSET: row.summary_conversation = _dumps(summary_conversation) if action_plan_thread_id is not _UNSET: row.action_plan_thread_id = action_plan_thread_id or None if action_plan_snapshot is not _UNSET: row.action_plan_snapshot = _dumps(action_plan_snapshot) if action_plan_conversation is not _UNSET: row.action_plan_conversation = _dumps(action_plan_conversation) if status is not _UNSET: row.status = str(status) row.version += 1 row.updated_at = datetime.now(UTC) await session.commit() await session.refresh(row) return self._session_dict(row) async def activate_session( self, session_id: str, user_id: str, *, chain_id: str, chain_snapshot: dict[str, Any], nodes: list[dict[str, Any]], ) -> tuple[dict[str, Any], list[dict[str, Any]]] | None: """Freeze a chain snapshot and create its nodes once. Repeated requests are intentionally idempotent: once nodes exist the existing snapshot and nodes are returned unchanged. """ now = datetime.now(UTC) async with self._sf() as session: row = await session.get(PositionRoundtableSessionRow, session_id) if row is None or not self._can_access(row, user_id): return None node_rows = ( await session.execute( select(PositionRoundtableNodeRow) .where(PositionRoundtableNodeRow.session_id == session_id) .order_by(PositionRoundtableNodeRow.stage_index, PositionRoundtableNodeRow.seat_index) ) ).scalars().all() if node_rows: return self._session_dict(row), [self._node_dict(node) for node in node_rows] row.chain_id = chain_id row.chain_snapshot = _dumps(chain_snapshot) row.status = "active" row.updated_at = now node_rows = [ PositionRoundtableNodeRow( id=uuid4().hex, session_id=session_id, node_key=str(node["node_key"]), stage_index=int(node["stage_index"]), seat_index=int(node["seat_index"]), agent_id=str(node["agent_id"]), position_id=node.get("position_id") or None, status=str(node.get("status") or "locked"), latest_answer=node.get("latest_answer") or None, artifact_manifest=_dumps(node.get("artifact_manifest") or []), revision=int(node.get("revision") or 0), upstream_revision=_dumps(node.get("upstream_revision")) if node.get("upstream_revision") is not None else None, created_at=now, updated_at=now, ) for node in nodes ] session.add_all(node_rows) await session.commit() await session.refresh(row) for node in node_rows: await session.refresh(node) return self._session_dict(row), [self._node_dict(node) for node in node_rows] async def list_nodes(self, session_id: str, user_id: str) -> list[dict[str, Any]] | None: async with self._sf() as session: parent = await session.get(PositionRoundtableSessionRow, session_id) if parent is None or not self._can_access(parent, user_id): return None rows = ( await session.execute( select(PositionRoundtableNodeRow) .where(PositionRoundtableNodeRow.session_id == session_id) .order_by(PositionRoundtableNodeRow.stage_index, PositionRoundtableNodeRow.seat_index) ) ).scalars().all() return [self._node_dict(row) for row in rows] async def get_node(self, session_id: str, node_key: str, user_id: str) -> dict[str, Any] | None: async with self._sf() as session: parent = await session.get(PositionRoundtableSessionRow, session_id) if parent is None or not self._can_access(parent, user_id): return None stmt = select(PositionRoundtableNodeRow).where( PositionRoundtableNodeRow.session_id == session_id, PositionRoundtableNodeRow.node_key == node_key, ).limit(1) row = (await session.execute(stmt)).scalars().first() return self._node_dict(row) if row is not None else None async def update_node( self, session_id: str, node_key: str, user_id: str, *, thread_id: str | None | object = _UNSET, status: str | object = _UNSET, latest_answer: str | None | object = _UNSET, artifact_manifest: list[dict[str, Any]] | object = _UNSET, conversation_snapshot: dict[str, Any] | None | object = _UNSET, upstream_revision: dict[str, Any] | object = _UNSET, last_rejection: dict[str, Any] | None | object = _UNSET, invalidated_by: str | None | object = _UNSET, expected_version: int | None = None, increment_revision: bool = False, increment_rejection: bool = False, ) -> dict[str, Any] | None: async with self._sf() as session: parent = await session.get(PositionRoundtableSessionRow, session_id) if parent is None or not self._can_access(parent, user_id): return None stmt = select(PositionRoundtableNodeRow).where( PositionRoundtableNodeRow.session_id == session_id, PositionRoundtableNodeRow.node_key == node_key, ).limit(1) # 乐观锁:传入 expected_version 时加行锁并校验,杜绝「检查后到提交之间 # 被并发改写」的 TOCTOU。recompute_chain_state 不走这里(它自管事务)。 if expected_version is not None: stmt = stmt.with_for_update() row = (await session.execute(stmt)).scalars().first() if row is None: return None if expected_version is not None and row.version != expected_version: raise NodeConcurrentWriteError( node_key, expected=expected_version, actual=row.version ) if thread_id is not _UNSET: row.thread_id = thread_id or None if status is not _UNSET: row.status = str(status) if latest_answer is not _UNSET: row.latest_answer = latest_answer if artifact_manifest is not _UNSET: row.artifact_manifest = _dumps(artifact_manifest) if conversation_snapshot is not _UNSET: row.conversation_snapshot = _dumps(conversation_snapshot) if upstream_revision is not _UNSET: row.upstream_revision = _dumps(upstream_revision) if last_rejection is not _UNSET: row.last_rejection = _dumps(last_rejection) if invalidated_by is not _UNSET: row.invalidated_by = invalidated_by or None # 任意字段变更都 bump version(乐观锁单调递增),与业务交付 revision 解耦。 row.version += 1 if increment_revision: row.revision += 1 if increment_rejection: row.rejection_count += 1 row.updated_at = datetime.now(UTC) await session.commit() await session.refresh(row) return self._node_dict(row) async def recompute_chain_state( self, session_id: str, node_key: str, user_id: str, *, next_status: str, latest_answer: str | None, artifact_manifest: list[dict[str, Any]] | None, upstream_revision: dict[str, Any] | None, new_effective_delivery: bool, ) -> dict[str, Any] | None: """在单一 ``FOR UPDATE`` 事务内原子完成「本节点交付 → 下游失效 → 解锁下一 stage」。 把原来 router 里由若干个独立提交的 ``update_node`` 拼成的状态转移收敛成一个 事务:先锁住本 session 全部节点行(MySQL ``SELECT ... FOR UPDATE``),再在锁 内重算所有节点状态,一次提交。同一 session 的并发 ``complete_node_turn`` 因此被 数据库行锁串行化,彻底消除 TOCTOU / 丢失更新导致的「下一 stage 卡 locked / invalidated_by 残留 → 下一阶段拿不到上游产物」断链。 幂等可重入:解锁不假设「最后一个 seat 负责」,而是遍历每个「全 done」stage 补 解锁其紧邻下一 stage 的 locked 节点;多次调用结果一致。``next_status`` / ``new_effective_delivery`` 等业务判定由 router 基于调用前快照算好传入,本方法 只负责把它们原子地落到节点状态机上。 """ async with self._sf() as session: parent = await session.get(PositionRoundtableSessionRow, session_id) if parent is None or not self._can_access(parent, user_id): return None # 锁住本 session 全部节点行,串行化并发状态转移(行锁持有至 commit)。 rows = ( await session.execute( select(PositionRoundtableNodeRow) .where(PositionRoundtableNodeRow.session_id == session_id) .order_by( PositionRoundtableNodeRow.stage_index, PositionRoundtableNodeRow.seat_index, ) .with_for_update() ) ).scalars().all() if not rows: return None by_key = {r.node_key: r for r in rows} current = by_key.get(node_key) if current is None: return None now = datetime.now(UTC) # 本次交付前的 revision(用于判断「此前已交付过 → 新交付需失效下游」)。 old_revision = current.revision # ① 落本节点交付状态。 current.status = next_status current.latest_answer = latest_answer if artifact_manifest is not None: current.artifact_manifest = _dumps(artifact_manifest) if upstream_revision is not None: current.upstream_revision = _dumps(upstream_revision) current.invalidated_by = None if new_effective_delivery: current.revision += 1 current.version += 1 current.updated_at = now # ② + ③ 与命令式状态机共用同一套下游失效 / 幂等解锁实现,避免逻辑漂移。 self._invalidate_downstream(rows, current, new_effective_delivery, old_revision, now) self._recompute_unlocks(rows, now) await session.commit() await session.refresh(current) return self._node_dict(current) # ── command-style state machine ────────────────────────────────────────── async def apply_node_command( self, session_id: str, node_key: str, user_id: str, *, command_id: str, command_type: str, payload: dict[str, Any] | None = None, expected_session_version: int | None = None, ) -> dict[str, Any]: """命令式状态机:所有节点级写命令(complete/bind/conversation/reject)的唯一入口。 单事务内依次完成: ① ``SELECT ... FOR UPDATE`` 锁住 session 行(可选 ``expected_session_version`` CAS,命中失败即 409); ② 命令幂等判定——``UNIQUE(session_id, node_key, command_id)`` 命中已有记录: 参数一致(payload_hash 相同)→ 原样重放上次结果;参数不同 → 409; ③ 锁住本 session 全部节点行,**基于锁内状态**做所有业务判定 (能否完成 / 是否新交付 / 能否驳回),绝不信任调用方传入的过期快照; ④ 应用状态转移(本节点更新 → 下游失效 → 幂等重算解锁); ⑤ 写入命令幂等记录(含结果 JSON),一次提交。 返回 dict: - ``ok=True`` — ``replayed`` 标记是否幂等重放;``node`` 为目标节点新状态, ``reject`` 命令额外返回 ``nodes``(全部节点)与 ``session``; - ``ok=False`` — ``http_status``(404/409)+ ``detail``;事务已回滚,状态未变。 一致性只来自数据库:并发命令被 session/节点行锁串行化,重试/双击/beacon 由 命令唯一键去重。进程内任何缓存都不参与正确性。 """ payload = payload or {} payload_hash = _payload_hash(command_type, payload) now = datetime.now(UTC) async def _run(session: AsyncSession) -> dict[str, Any]: parent = ( await session.execute( select(PositionRoundtableSessionRow) .where(PositionRoundtableSessionRow.id == session_id) .with_for_update() ) ).scalar_one_or_none() if parent is None or not self._can_access(parent, user_id): return {"ok": False, "http_status": 404, "detail": "Position roundtable session not found"} if expected_session_version is not None and parent.version != expected_session_version: return { "ok": False, "http_status": 409, "detail": "会话状态已更新,请刷新后重试", "session": self._session_dict(parent), } # 归档后的会话拒绝一切写命令(锁内复检,杜绝路由预检与提交之间的窗口)。 if parent.status == "archived": return { "ok": False, "http_status": 409, "detail": "该岗位协同会话已归档,仅可查看历史记录", } # ② 命令幂等:同 command_id 已执行过 → 重放上次结果(或 409)。 existing = ( await session.execute( select(PositionRoundtableCommandRow).where( PositionRoundtableCommandRow.session_id == session_id, PositionRoundtableCommandRow.node_key == (node_key or ""), PositionRoundtableCommandRow.command_id == command_id, ) ) ).scalars().first() if existing is not None: if existing.payload_hash != payload_hash: return { "ok": False, "http_status": 409, "detail": "相同操作标识携带了不同参数,已拒绝执行", } return { "ok": True, "replayed": True, "command_type": command_type, **(_loads(existing.result_json) or {}), } # ③ 锁住全部节点行(行锁持有至 commit),串行化同 session 并发命令。 rows = ( await session.execute( select(PositionRoundtableNodeRow) .where(PositionRoundtableNodeRow.session_id == session_id) .order_by( PositionRoundtableNodeRow.stage_index, PositionRoundtableNodeRow.seat_index, ) .with_for_update() ) ).scalars().all() by_key = {r.node_key: r for r in rows} if command_type == "complete": result = self._apply_complete(by_key, node_key, parent, payload, now) elif command_type == "bind": result = self._apply_bind(by_key, node_key, payload, now) elif command_type == "conversation": result = self._apply_conversation(by_key, node_key, payload, now) elif command_type == "reject": result = self._apply_reject(by_key, node_key, parent, payload, user_id, now) else: return {"ok": False, "http_status": 400, "detail": f"unknown command type: {command_type}"} if not result.get("ok", True): return result # ⑤ 写命令幂等记录(与状态转移同事务提交)。 stored = {k: v for k, v in result.items() if k != "ok"} session.add( PositionRoundtableCommandRow( id=uuid4().hex, session_id=session_id, node_key=node_key or "", command_id=command_id, command_type=command_type, payload_hash=payload_hash, result_json=_dumps(stored), created_at=now, ) ) parent.version += 1 parent.updated_at = now result["replayed"] = False result["command_type"] = command_type return result try: async with self._sf() as session: outcome = await _run(session) if outcome.get("ok", False): await session.commit() else: await session.rollback() return outcome except IntegrityError: # 命令唯一键并发竞态:另一连接已用同一 command_id 提交。读回其结果重放。 async with self._sf() as session: existing = ( await session.execute( select(PositionRoundtableCommandRow).where( PositionRoundtableCommandRow.session_id == session_id, PositionRoundtableCommandRow.node_key == (node_key or ""), PositionRoundtableCommandRow.command_id == command_id, ) ) ).scalars().first() if existing is not None: if existing.payload_hash != payload_hash: return { "ok": False, "http_status": 409, "detail": "相同操作标识携带了不同参数,已拒绝执行", } return { "ok": True, "replayed": True, "command_type": command_type, **(_loads(existing.result_json) or {}), } raise # ── command handlers(全部在 FOR UPDATE 锁内调用,入参为锁定的行) ───────── def _apply_complete( self, by_key: dict[str, PositionRoundtableNodeRow], node_key: str, parent: PositionRoundtableSessionRow, payload: dict[str, Any], now: datetime, ) -> dict[str, Any]: """complete-turn:锁内判定能否完成 / 是否新交付,原子落状态机。""" current = by_key.get(node_key) if current is None: return {"ok": False, "http_status": 404, "detail": "Position roundtable node not found"} if current.status == "locked": return {"ok": False, "http_status": 409, "detail": "上游节点尚未完成,当前节点未解锁"} if current.invalidated_by: return { "ok": False, "http_status": 409, "detail": "上游产物已被驳回或重做,当前节点需等待前序节点重新完成", } # 请求未带产物清单 → 沿用锁内的当前清单(绝不用调用方的过期快照兜底)。 incoming = payload.get("artifact_manifest") current_manifest = _loads(current.artifact_manifest) or [] effective_artifacts = incoming if isinstance(incoming, list) and incoming else current_manifest latest_answer = payload.get("latest_answer") upstream_revision = payload.get("upstream_revision") has_markdown = any(is_markdown_artifact(item) for item in effective_artifacts) markdown_changed = markdown_delivery_changed(current_manifest, effective_artifacts) requires_fresh_delivery = current.status in {"rejected", "stale"} intent_position_id = session_intent_position_id(_loads(parent.chain_snapshot)) is_intent_node = is_intent_position_node(current.position_id, intent_position_id) next_status = ( "done" if is_intent_node else "done" if has_markdown and (not requires_fresh_delivery or markdown_changed) else "running" ) # 有效 revision 是「交付 revision」,不是每一次对话应答。已完成节点上的普通 # 追问,不能仅因其旧 Markdown 仍出现在清单里就失效所有下游 stage。 new_effective_delivery = ( next_status == "done" and (is_intent_node or current.status != "done" or markdown_changed) ) old_revision = current.revision # ④ 落本节点交付状态。 current.status = next_status current.latest_answer = latest_answer current.artifact_manifest = _dumps(effective_artifacts) if upstream_revision is not None: current.upstream_revision = _dumps(upstream_revision) current.invalidated_by = None if new_effective_delivery: current.revision += 1 current.version += 1 current.updated_at = now self._invalidate_downstream(by_key.values(), current, new_effective_delivery, old_revision, now) self._recompute_unlocks(by_key.values(), now) return {"ok": True, "node": self._node_dict(current)} def _apply_bind( self, by_key: dict[str, PositionRoundtableNodeRow], node_key: str, payload: dict[str, Any], now: datetime, ) -> dict[str, Any]: """bind-thread:锁内校验解锁/失效/线程占用后绑定。""" current = by_key.get(node_key) if current is None: return {"ok": False, "http_status": 404, "detail": "Position roundtable node not found"} thread_id = str(payload.get("thread_id") or "").strip() if not thread_id: return {"ok": False, "http_status": 400, "detail": "thread_id is required"} if current.status == "locked": return {"ok": False, "http_status": 409, "detail": "上游节点尚未完成,当前节点未解锁"} if current.invalidated_by: return { "ok": False, "http_status": 409, "detail": "上游产物已被驳回或重做,当前节点需等待前序节点重新完成", } if current.thread_id and current.thread_id != thread_id: return {"ok": False, "http_status": 409, "detail": "该节点已绑定另一条对话线程"} current.thread_id = thread_id if current.status == "ready": current.status = "running" current.version += 1 current.updated_at = now return {"ok": True, "node": self._node_dict(current)} def _apply_conversation( self, by_key: dict[str, PositionRoundtableNodeRow], node_key: str, payload: dict[str, Any], now: datetime, ) -> dict[str, Any]: """conversation:对话快照持久化(UI 恢复用),锁内校验线程归属。""" current = by_key.get(node_key) if current is None: return {"ok": False, "http_status": 404, "detail": "Position roundtable node not found"} thread_id = payload.get("thread_id") if thread_id and current.thread_id and current.thread_id != thread_id: return {"ok": False, "http_status": 409, "detail": "Position node is bound to another thread"} if thread_id and not current.thread_id: current.thread_id = thread_id current.conversation_snapshot = _dumps(payload.get("conversation_snapshot")) current.version += 1 current.updated_at = now return {"ok": True, "node": self._node_dict(current)} def _apply_reject( self, by_key: dict[str, PositionRoundtableNodeRow], node_key: str, parent: PositionRoundtableSessionRow, payload: dict[str, Any], user_id: str, now: datetime, ) -> dict[str, Any]: """reject:驳回一个已交付节点并失效全部下游 stage(锁内判定可驳回性)。""" rows = sorted( by_key.values(), key=lambda r: (r.stage_index, r.seat_index) ) target = by_key.get(node_key) if target is None: return {"ok": False, "http_status": 404, "detail": "Position roundtable node not found"} if target.status == "locked": return {"ok": False, "http_status": 409, "detail": "当前节点尚未解锁,暂无可驳回的交付产物"} if ( target.revision <= 0 and not target.latest_answer and not _loads(target.artifact_manifest) ): return {"ok": False, "http_status": 409, "detail": "当前节点暂无可驳回的交付产物"} rejection = { "title": str(payload.get("title") or "").strip() or "产物驳回", "reason": str(payload.get("reason") or "").strip(), "rejected_by": user_id, "rejected_at": now.isoformat(), } summary_snapshot = _loads(parent.summary_snapshot) if isinstance(summary_snapshot, dict) and summary_snapshot.get("status") == "done": parent.summary_snapshot = _dumps( { **summary_snapshot, "status": "stale", "invalidated_by": target.node_key, "invalidated_at": now.isoformat(), } ) action_plan_snapshot = _loads(parent.action_plan_snapshot) if isinstance(action_plan_snapshot, dict) and action_plan_snapshot.get("status") == "done": parent.action_plan_snapshot = _dumps( { **action_plan_snapshot, "status": "stale", "invalidated_by": target.node_key, "invalidated_at": now.isoformat(), } ) if parent.status == "completed": parent.status = "active" target.status = "rejected" target.last_rejection = _dumps(rejection) target.rejection_count += 1 target.invalidated_by = None target.version += 1 target.updated_at = now for row in rows: if row.stage_index <= target.stage_index: continue if _has_material(row) or row.status in {"running", "done", "stale", "rejected", "error"}: row.status = "stale" else: row.status = "locked" row.invalidated_by = target.node_key row.version += 1 row.updated_at = now return { "ok": True, "session": self._session_dict(parent), "nodes": [self._node_dict(row) for row in rows], } @staticmethod def _invalidate_downstream( rows: Any, current: PositionRoundtableNodeRow, new_effective_delivery: bool, old_revision: int, now: datetime, ) -> None: """新交付且本节点此前已交付过 → 后续 stage 全部失效(需重做)。""" if not (new_effective_delivery and old_revision > 0): return for downstream in rows: if ( downstream.stage_index > current.stage_index and downstream.status in {"ready", "running", "done", "stale", "rejected", "error"} ): downstream.status = "stale" if _has_material(downstream) else "locked" downstream.invalidated_by = current.node_key downstream.version += 1 downstream.updated_at = now @staticmethod def _recompute_unlocks(rows: Any, now: datetime) -> None: """幂等重算 stage 解锁:每个「全部 done」stage 解锁其紧邻下一 stage。 locked → ready;stale 保持 stale(产物过期需重做)但清掉 invalidated_by (上游已就绪)。多次调用结果一致,不依赖调用时序。 """ row_list = list(rows) stage_indices = sorted({r.stage_index for r in row_list}) for stage in stage_indices: stage_rows = [r for r in row_list if r.stage_index == stage] if not stage_rows or not all(r.status == "done" for r in stage_rows): continue later = [t for t in stage_indices if t > stage] if not later: continue next_stage = min(later) for node in row_list: if node.stage_index != next_stage: continue if node.status == "locked": node.status = "ready" node.invalidated_by = None node.version += 1 node.updated_at = now elif node.status == "stale": node.invalidated_by = None node.version += 1 node.updated_at = now async def apply_session_command( self, session_id: str, user_id: str, *, command_id: str, command_type: str, payload: dict[str, Any] | None = None, expected_session_version: int | None = None, ) -> dict[str, Any]: """session 级命令(activate / archive / delete)的命令式入口。 与 ``apply_node_command`` 同一套原则:单事务锁 session 行(``FOR UPDATE``) + 命令唯一键幂等 + 锁内判定。返回结构同 ``apply_node_command`` (``ok`` / ``replayed`` / ``http_status`` / ``detail`` / ``session`` / ``nodes``)。 - ``activate`` — payload 需带 ``chain_id`` / ``chain_snapshot`` / ``nodes``。 锁内判定:已激活且链相同 → 幂等返回现有状态;链不同 → 409。 - ``archive`` — 已归档 → 幂等返回;否则置 archived。 - ``delete`` — 删除节点与会话(命令记录无外键,保留供重放)。 """ payload = payload or {} payload_hash = _payload_hash(command_type, payload) now = datetime.now(UTC) async def _replay_check(session: AsyncSession) -> dict[str, Any] | None: existing = ( await session.execute( select(PositionRoundtableCommandRow).where( PositionRoundtableCommandRow.session_id == session_id, PositionRoundtableCommandRow.node_key == "", PositionRoundtableCommandRow.command_id == command_id, ) ) ).scalars().first() if existing is None: return None if existing.payload_hash != payload_hash: return { "ok": False, "http_status": 409, "detail": "相同操作标识携带了不同参数,已拒绝执行", } return { "ok": True, "replayed": True, "command_type": command_type, **(_loads(existing.result_json) or {}), } async def _run(session: AsyncSession) -> dict[str, Any]: # delete 之后 session 行已不存在,重试必须在锁 session 之前查命令记录。 replayed = await _replay_check(session) if replayed is not None: return replayed parent = ( await session.execute( select(PositionRoundtableSessionRow) .where(PositionRoundtableSessionRow.id == session_id) .with_for_update() ) ).scalar_one_or_none() if parent is None or not self._can_access(parent, user_id): return {"ok": False, "http_status": 404, "detail": "Position roundtable session not found"} if expected_session_version is not None and parent.version != expected_session_version: return { "ok": False, "http_status": 409, "detail": "会话状态已更新,请刷新后重试", "session": self._session_dict(parent), } if command_type == "activate": result = await self._apply_activate(session, parent, payload, now) elif command_type == "archive": if parent.status == "archived": result = { "ok": True, "session": self._session_dict(parent), "nodes": [], "already": True, } else: parent.status = "archived" parent.version += 1 parent.updated_at = now result = { "ok": True, "session": self._session_dict(parent), "nodes": [], } elif command_type == "delete": await session.execute( delete(PositionRoundtableNodeRow).where( PositionRoundtableNodeRow.session_id == session_id ) ) await session.execute( delete(PositionRoundtableSessionRow).where( PositionRoundtableSessionRow.id == session_id ) ) result = {"ok": True, "deleted": True} else: return {"ok": False, "http_status": 400, "detail": f"unknown command type: {command_type}"} if not result.get("ok", True): return result stored = {k: v for k, v in result.items() if k != "ok"} session.add( PositionRoundtableCommandRow( id=uuid4().hex, session_id=session_id, node_key="", command_id=command_id, command_type=command_type, payload_hash=payload_hash, result_json=_dumps(stored), created_at=now, ) ) result["replayed"] = False result["command_type"] = command_type return result try: async with self._sf() as session: outcome = await _run(session) if outcome.get("ok", False): await session.commit() else: await session.rollback() return outcome except IntegrityError: # 命令唯一键并发竞态:读回已提交的那条命令结果重放。 async with self._sf() as session: replayed = await _replay_check(session) if replayed is not None: return replayed raise async def _apply_activate( self, session: AsyncSession, parent: PositionRoundtableSessionRow, payload: dict[str, Any], now: datetime, ) -> dict[str, Any]: """激活业务链:锁内判定「已激活 → 链相同幂等 / 链不同 409」,否则建节点。""" existing_rows = ( await session.execute( select(PositionRoundtableNodeRow) .where(PositionRoundtableNodeRow.session_id == parent.id) .order_by( PositionRoundtableNodeRow.stage_index, PositionRoundtableNodeRow.seat_index, ) .with_for_update() ) ).scalars().all() if parent.status == "archived": return { "ok": False, "http_status": 409, "detail": "该岗位协同会话已归档,仅可查看历史记录", } chain_id = str(payload.get("chain_id") or "") if existing_rows: # 已激活:同一业务链 → 幂等返回现有状态;不同链 → 冲突。 if parent.chain_id == chain_id: return { "ok": True, "session": self._session_dict(parent), "nodes": [self._node_dict(row) for row in existing_rows], "already": True, } return { "ok": False, "http_status": 409, "detail": "该会话已激活其他业务链条,不能重复激活", } if not _has_confirmed_intent_snapshot(_loads(parent.intent_snapshot)): return { "ok": False, "http_status": 400, "detail": "请先完成并保存任务意图", } chain_snapshot = payload.get("chain_snapshot") or {} nodes = payload.get("nodes") or [] if not nodes: return {"ok": False, "http_status": 400, "detail": "activate requires nodes"} parent.chain_id = chain_id parent.chain_snapshot = _dumps(chain_snapshot) parent.status = "active" parent.version += 1 parent.updated_at = now node_rows = [ PositionRoundtableNodeRow( id=uuid4().hex, session_id=parent.id, node_key=str(node["node_key"]), stage_index=int(node["stage_index"]), seat_index=int(node["seat_index"]), agent_id=str(node["agent_id"]), position_id=node.get("position_id") or None, status=str(node.get("status") or "locked"), latest_answer=node.get("latest_answer") or None, artifact_manifest=_dumps(node.get("artifact_manifest") or []), revision=int(node.get("revision") or 0), upstream_revision=( _dumps(node.get("upstream_revision")) if node.get("upstream_revision") is not None else None ), created_at=now, updated_at=now, ) for node in nodes ] session.add_all(node_rows) # 同事务内 flush 拿到完整行状态,供命令记录与返回值使用。 await session.flush() return { "ok": True, "session": self._session_dict(parent), "nodes": [self._node_dict(row) for row in node_rows], } async def reject_node( self, session_id: str, node_key: str, user_id: str, *, title: str, reason: str, rejected_by: str, ) -> tuple[dict[str, Any], list[dict[str, Any]]] | None: """Reject one delivered node and invalidate every later stage. Historical answers and artifact manifests are deliberately preserved. ``invalidated_by`` gates downstream rework until the rejected upstream stage is completed again. """ now = datetime.now(UTC) async with self._sf() as session: parent = await session.get(PositionRoundtableSessionRow, session_id) if parent is None or not self._can_access(parent, user_id): return None rows = ( await session.execute( select(PositionRoundtableNodeRow) .where(PositionRoundtableNodeRow.session_id == session_id) .order_by(PositionRoundtableNodeRow.stage_index, PositionRoundtableNodeRow.seat_index) ) ).scalars().all() target = next((row for row in rows if row.node_key == node_key), None) if target is None: return None rejection = { "title": title, "reason": reason, "rejected_by": rejected_by, "rejected_at": now.isoformat(), } summary_snapshot = _loads(parent.summary_snapshot) if isinstance(summary_snapshot, dict) and summary_snapshot.get("status") == "done": summary_snapshot = { **summary_snapshot, "status": "stale", "invalidated_by": target.node_key, "invalidated_at": now.isoformat(), } parent.summary_snapshot = _dumps(summary_snapshot) action_plan_snapshot = _loads(parent.action_plan_snapshot) if isinstance(action_plan_snapshot, dict) and action_plan_snapshot.get("status") == "done": action_plan_snapshot = { **action_plan_snapshot, "status": "stale", "invalidated_by": target.node_key, "invalidated_at": now.isoformat(), } parent.action_plan_snapshot = _dumps(action_plan_snapshot) if parent.status == "completed": parent.status = "active" target.status = "rejected" target.last_rejection = _dumps(rejection) target.rejection_count += 1 target.invalidated_by = None target.updated_at = now for row in rows: if row.stage_index <= target.stage_index: continue if _has_material(row) or row.status in {"running", "done", "stale", "rejected", "error"}: row.status = "stale" else: row.status = "locked" row.invalidated_by = target.node_key row.updated_at = now parent.updated_at = now await session.commit() await session.refresh(parent) for row in rows: await session.refresh(row) return self._session_dict(parent), [self._node_dict(row) for row in rows] async def archive_session(self, session_id: str, user_id: str) -> dict[str, Any] | None: return await self.update_session(session_id, user_id, status="archived") async def delete_session(self, session_id: str, user_id: str) -> bool: """Delete a task-shared row or the caller's personal legacy row.""" async with self._sf() as session: parent = ( await session.execute( select(PositionRoundtableSessionRow) .where( PositionRoundtableSessionRow.id == session_id, ) .limit(1) ) ).scalar_one_or_none() if parent is None or not self._can_access(parent, user_id): return False # Do not rely on database-specific FK cascade settings. SQLite # deployments may not have PRAGMA foreign_keys enabled, which # would otherwise leave orphaned conversation/artifact snapshots. await session.execute( delete(PositionRoundtableNodeRow).where( PositionRoundtableNodeRow.session_id == session_id ) ) await session.execute( delete(PositionRoundtableSessionRow).where( PositionRoundtableSessionRow.id == session_id, ) ) await session.commit() return True async def list_results(self, session_id: str, user_id: str) -> list[dict[str, Any]] | None: return await self.list_nodes(session_id, user_id)