deerflow-code/offline-backend-20260512/backend/packages/harness/deerflow/persistence/intent_turns/sql.py
2026-09-07 18:24:55 +08:00

226 lines
8.5 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""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