226 lines
8.5 KiB
Python
226 lines
8.5 KiB
Python
"""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
|