1435 lines
62 KiB
Python
1435 lines
62 KiB
Python
"""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)
|