"""SQLAlchemy-backed intent-turn single-flight storage (Phase 4). 跨 worker 单飞 + 幂等回放,一致性只来自 DB(唯一键 + 条件 UPDATE),不依赖进程内 ``RunManager``(多 worker 下彼此不可见)。两条核心路径: * ``try_claim`` —— INSERT 占用 ``(thread_id, client_turn_id)``;唯一键冲突即幂等命中, 按「已 done→回放 / 在跑且租约有效→丢弃 / 租约过期或 error→抢占续跑 / 内容不一致→409」分流。 * ``record_result`` —— 持租者把终态帧写回(条件 UPDATE 要求租约匹配),清租约。 """ from __future__ import annotations import json from datetime import UTC, datetime from typing import Any from uuid import uuid4 from sqlalchemy import select, update from sqlalchemy.exc import IntegrityError from sqlalchemy.ext.asyncio import async_sessionmaker from deerflow.persistence.intent_turns.model import IntentTurnRow def _loads(value: Any) -> Any: if isinstance(value, str) and value: try: return json.loads(value) except (json.JSONDecodeError, TypeError): return value return value class IntentTurnRepository: def __init__(self, session_factory: async_sessionmaker) -> None: self._sf = session_factory @staticmethod def _row_to_dict(row: IntentTurnRow | None) -> dict[str, Any] | None: if row is None: return None return { "id": row.id, "thread_id": row.thread_id, "client_turn_id": row.client_turn_id, "message_hash": row.message_hash, "status": row.status, "result_json": _loads(row.result_json), "error": row.error, "lease_owner": row.lease_owner, "lease_until": row.lease_until.isoformat() if row.lease_until else None, "created_at": row.created_at, "updated_at": row.updated_at, } async def get(self, thread_id: str, client_turn_id: str) -> dict[str, Any] | None: stmt = select(IntentTurnRow).where( IntentTurnRow.thread_id == thread_id, IntentTurnRow.client_turn_id == client_turn_id, ) async with self._sf() as session: row = (await session.execute(stmt)).scalar_one_or_none() return self._row_to_dict(row) async def try_claim( self, *, thread_id: str, client_turn_id: str, message_hash: str, lease_owner: str, lease_until: datetime, ) -> dict[str, Any]: """占用一个意图回合 —— 跨 worker 单飞 + 幂等的入口。 返回 ``{"action": ..., "row": {...}}``,``action`` 取值: * ``"run"`` —— 本调用方赢得占用(新建 / 抢占续跑),应启 run 并在终态写回。 * ``"replay"`` —— 该回合已 ``done``,原样回放 ``row.result_json``,**不再启 run**。 * ``"duplicate"``—— 该回合仍在跑且租约有效(或刚被别处抢占),本请求是重复的 in-flight 请求,应丢弃(原请求仍在流式吐结果)。 * ``"conflict"`` —— 同一 ``client_turn_id`` 但 ``message_hash`` 不同 → 客户端 bug,409。 """ # 1) 先尝试 INSERT 占用。 new_id = uuid4().hex try: async with self._sf() as session: session.add( IntentTurnRow( id=new_id, thread_id=thread_id, client_turn_id=client_turn_id, message_hash=message_hash, status="running", lease_owner=lease_owner, lease_until=lease_until, ) ) await session.commit() row = await self.get(thread_id, client_turn_id) return {"action": "run", "row": row} except IntegrityError: # 唯一键 (thread_id, client_turn_id) 冲突 → 幂等命中,读回既有行分流。 pass existing = await self.get(thread_id, client_turn_id) if existing is None: # 极端竞态:冲突与回读之间行被清掉。降级为「可新建」再插一次(best-effort)。 return {"action": "run", "row": None} # 2) 内容指纹不一致 → 客户端 bug(同 client_turn_id 却换了 message)。 if existing.get("message_hash") != message_hash: return {"action": "conflict", "row": existing} status = existing.get("status") or "running" # 3) 已完成 → 幂等回放终态帧。 if status == "done": return {"action": "replay", "row": existing} # 4) 在跑或出错 → 尝试条件 UPDATE 抢占续跑(error 强制;running 要求租约已过期)。 reclaimed = await self._conditional_reclaim( thread_id=thread_id, client_turn_id=client_turn_id, prev_status=status, lease_owner=lease_owner, lease_until=lease_until, ) if reclaimed is not None: return {"action": "run", "row": reclaimed} # 5) 抢占失败:仍在跑且租约有效(或刚被别处抢占)→ 本请求是重复 in-flight,丢弃。 return {"action": "duplicate", "row": existing} async def _conditional_reclaim( self, *, thread_id: str, client_turn_id: str, prev_status: str, lease_owner: str, lease_until: datetime, ) -> dict[str, Any] | None: """原子条件 UPDATE 抢占:``error`` 强制续跑;``running`` 仅当租约已过期。 成功则置 ``status='running'``、写新租约,返回行;否则 None(仍在跑 / 刚被抢占)。 """ now = datetime.now(UTC) where = [ IntentTurnRow.thread_id == thread_id, IntentTurnRow.client_turn_id == client_turn_id, ] if prev_status == "error": where.append(IntentTurnRow.status == "error") else: where.extend( [ IntentTurnRow.status == "running", IntentTurnRow.lease_until.is_not(None), IntentTurnRow.lease_until < now, ] ) values = { "status": "running", "lease_owner": lease_owner, "lease_until": lease_until, "updated_at": now, } async with self._sf() as session: result = await session.execute( update(IntentTurnRow).where(*where).values(**values) ) if result.rowcount == 0: return None await session.commit() row = ( await session.execute( select(IntentTurnRow).where( IntentTurnRow.thread_id == thread_id, IntentTurnRow.client_turn_id == client_turn_id, ) ) ).scalar_one_or_none() return self._row_to_dict(row) async def record_result( self, *, thread_id: str, client_turn_id: str, lease_owner: str, status: str, result_json: Any = None, error: str | None = None, ) -> bool: """持租者把终态帧写回(条件 UPDATE 要求租约匹配),并清租约。 只有赢得 ``try_claim``(``action='run'``)的调用方才应调用,且 ``lease_owner`` 必须与领取时一致 —— 防止「抢占续跑的新 worker」被「旧 worker 的迟到收尾」覆盖。 """ now = datetime.now(UTC) values: dict[str, Any] = { "status": status, "lease_owner": None, "lease_until": None, "updated_at": now, } if result_json is not None: values["result_json"] = json.dumps(result_json, ensure_ascii=False) if error is not None: values["error"] = error where = [ IntentTurnRow.thread_id == thread_id, IntentTurnRow.client_turn_id == client_turn_id, IntentTurnRow.lease_owner == lease_owner, IntentTurnRow.status == "running", ] async with self._sf() as session: result = await session.execute( update(IntentTurnRow).where(*where).values(**values) ) if result.rowcount == 0: return False await session.commit() return True