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

285 lines
12 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 roundtable-planning draft + recommendation-history store.
JSON columns (``step1`` / ``step2`` / ``picks`` / ``candidates``) are stored as
serialized strings in ``PortableLongText`` and round-tripped here with
``json.dumps`` / ``json.loads`` — same approach as ``ai_writing_sessions`` — so
large Step 2 transcripts never hit MySQL's 64 KB ``TEXT`` cap.
All write/read paths are scoped by ``user_id`` for ownership isolation.
"""
from __future__ import annotations
import json
import re
from datetime import UTC, datetime
from typing import Any
from sqlalchemy import case, select, update
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
from deerflow.persistence.roundtable_drafts.model import (
RoundtableDraftRow,
RoundtableRecommendHistoryRow,
)
_UNSET = object()
# 业务链名抽取:在已加载进内存的 step2 JSON 串里直接找 `"chain":{...,"title":"X"}`,
# 不整体 json.loads(step2 是 LONGTEXT 大块),只做一次廉价正则扫描。`"chain":null`
# (recommend/自由模式)不匹配 → 返回 None。`[^{}]*?` 限定在 chain 对象内({id,title}
# 无嵌套花括号)。
_CHAIN_TITLE_RE = re.compile(r'"chain"\s*:\s*\{[^{}]*?"title"\s*:\s*"((?:\\.|[^"\\])*)"')
def _extract_chain_title(step2_raw: Any) -> str | None:
"""从 step2 原始 JSON 串里抽取所用业务链名(无链条 → None,前端显示「自由模式」)。"""
if not isinstance(step2_raw, str) or not step2_raw:
return None
m = _CHAIN_TITLE_RE.search(step2_raw)
if not m:
return None
title = m.group(1)
try:
# 还原 JSON 字符串转义(\" \\ \n \uXXXX 等)。
title = json.loads(f'"{title}"')
except Exception:
pass
title = title.strip()
return title or None
def _loads(value: Any) -> Any:
if not isinstance(value, str):
return value
if not value:
return None
try:
return json.loads(value)
except Exception:
return None
def _dumps(value: Any) -> str | None:
if value is None:
return None
return json.dumps(value, ensure_ascii=False)
class DraftConcurrentWriteError(Exception):
"""乐观锁冲突:``update_draft(expected_version=...)`` 时版本号不匹配。
由后台作业 ``_write_run_to_draft`` 的读-改-写重试循环捕获:冲突时重新读取最新草稿、
重新合并 ``step2.runs[]``、再试。防多用户/多 worker 并发写同一草稿时后写覆盖先写。
"""
def __init__(self, draft_id: str, *, expected: int | None = None, actual: int | None = None) -> None:
self.draft_id = draft_id
self.expected = expected
self.actual = actual
super().__init__(
f"concurrent write detected on roundtable draft {draft_id} "
f"(expected version {expected}, got {actual})"
)
class RoundtableDraftRepository:
def __init__(self, session_factory: async_sessionmaker[AsyncSession]) -> None:
self._sf = session_factory
# ── drafts ────────────────────────────────────────────────────────────
@staticmethod
def _draft_to_meta(row: RoundtableDraftRow) -> dict[str, Any]:
"""Lightweight projection for the history dropdown (no big blobs)."""
return {
"id": row.id,
"task_id": row.task_id,
"title": row.title,
"furthest_step": row.furthest_step,
# 乐观锁版本号:列表也携带,供前端本地副本对账(冲突时重拉合并)。
"version": row.version,
# 所用业务链名(历史卡片展示用);无链条 → None → 前端显示「自由模式」。
"chain_title": _extract_chain_title(row.step2),
"created_at": row.created_at.isoformat() if isinstance(row.created_at, datetime) else row.created_at,
"updated_at": row.updated_at.isoformat() if isinstance(row.updated_at, datetime) else row.updated_at,
}
@classmethod
def _draft_to_dict(cls, row: RoundtableDraftRow) -> dict[str, Any]:
d = cls._draft_to_meta(row)
d["step1"] = _loads(row.step1)
d["step2"] = _loads(row.step2)
d["step3"] = _loads(row.step3)
return d
async def list_drafts(self, user_id: str, *, limit: int = 100) -> list[dict[str, Any]]:
stmt = (
select(RoundtableDraftRow)
.where(RoundtableDraftRow.user_id == user_id)
.order_by(RoundtableDraftRow.updated_at.desc())
.limit(limit)
)
async with self._sf() as session:
result = await session.execute(stmt)
return [self._draft_to_meta(row) for row in result.scalars()]
async def get_draft(self, draft_id: str, user_id: str) -> dict[str, Any] | None:
async with self._sf() as session:
row = await session.get(RoundtableDraftRow, draft_id)
# 个人草稿严格按 user 隔离。taskId 深链聊天记录已迁出到独立的
# ``roundtable_task_drafts`` 表(不分权),不再走这里。
if row is None or row.user_id != user_id:
return None
return self._draft_to_dict(row)
async def create_draft(self, user_id: str, data: dict[str, Any]) -> dict[str, Any]:
now = datetime.now(UTC)
row = RoundtableDraftRow(
id=data["id"],
user_id=user_id,
task_id=(data.get("task_id") or None),
title=(data.get("title") or "").strip(),
furthest_step=int(data.get("furthest_step") or 1),
step1=_dumps(data.get("step1")),
step2=_dumps(data.get("step2")),
step3=_dumps(data.get("step3")),
created_at=now,
updated_at=now,
)
async with self._sf() as session:
session.add(row)
await session.commit()
await session.refresh(row)
return self._draft_to_dict(row)
async def update_draft(
self,
draft_id: str,
user_id: str,
*,
title=_UNSET,
furthest_step=_UNSET,
step1=_UNSET,
step2=_UNSET,
step3=_UNSET,
task_id=_UNSET,
expected_version: int | None = None,
) -> dict[str, Any] | None:
"""原子条件 UPDATE:版本判定与字段写入发生在**同一条 SQL** 里。
一致性只来自数据库:``UPDATE ... WHERE id AND user_id AND version=`` 配合
rowcount 判定,多 worker 并发下不存在「读-改-写」的检查后竞争窗口。
``rowcount == 0`` 时再回读区分「不存在/非本人」(返回 None)与
「版本冲突」(抛 :class:`DraftConcurrentWriteError`)。
"""
values: dict[str, Any] = {
"version": RoundtableDraftRow.version + 1,
"updated_at": datetime.now(UTC),
}
if task_id is not _UNSET and task_id is not None:
values["task_id"] = str(task_id)
if title is not _UNSET and title is not None:
values["title"] = str(title).strip()
if furthest_step is not _UNSET and furthest_step is not None:
# furthest_step 单调不回退:迟到的自动保存不得把已推进的步骤拉回去。
step_value = int(furthest_step)
values["furthest_step"] = case(
(RoundtableDraftRow.furthest_step >= step_value, RoundtableDraftRow.furthest_step),
else_=step_value,
)
if step1 is not _UNSET:
values["step1"] = _dumps(step1)
if step2 is not _UNSET:
values["step2"] = _dumps(step2)
if step3 is not _UNSET:
values["step3"] = _dumps(step3)
where = [
RoundtableDraftRow.id == draft_id,
# 个人草稿严格按 user 隔离(task 聊天记录已迁出到独立存储)。
RoundtableDraftRow.user_id == user_id,
]
if expected_version is not None:
where.append(RoundtableDraftRow.version == expected_version)
async with self._sf() as session:
result = await session.execute(
update(RoundtableDraftRow).where(*where).values(**values)
)
if result.rowcount == 0:
if expected_version is not None:
# 回读区分冲突与不存在(回读在冲突分支,正常路径零额外开销)。
existing = await session.get(RoundtableDraftRow, draft_id)
if existing is not None and existing.user_id == user_id:
raise DraftConcurrentWriteError(
draft_id, expected=expected_version, actual=existing.version
)
return None
await session.commit()
row = await session.get(RoundtableDraftRow, draft_id)
await session.refresh(row)
return self._draft_to_dict(row)
async def delete_draft(self, draft_id: str, user_id: str) -> bool:
async with self._sf() as session:
row = await session.get(RoundtableDraftRow, draft_id)
# 个人草稿严格按 user 隔离(task 聊天记录已迁出到独立存储)。
if row is None or row.user_id != user_id:
return False
await session.delete(row)
await session.commit()
return True
# ── recommendation history ────────────────────────────────────────────
@staticmethod
def _history_to_dict(row: RoundtableRecommendHistoryRow) -> dict[str, Any]:
return {
"id": row.id,
"draft_id": row.draft_id,
"objective": row.objective,
"status": row.status,
"model": row.model,
"rationale": row.rationale,
"picks": _loads(row.picks) or [],
"candidates": _loads(row.candidates) or [],
"created_at": row.created_at.isoformat() if isinstance(row.created_at, datetime) else row.created_at,
}
async def list_recommendations(
self, draft_id: str, user_id: str, *, limit: int = 20
) -> list[dict[str, Any]]:
stmt = (
select(RoundtableRecommendHistoryRow)
.where(
RoundtableRecommendHistoryRow.draft_id == draft_id,
RoundtableRecommendHistoryRow.user_id == user_id,
)
.order_by(RoundtableRecommendHistoryRow.created_at.desc())
.limit(limit)
)
async with self._sf() as session:
result = await session.execute(stmt)
return [self._history_to_dict(row) for row in result.scalars()]
async def add_recommendation(self, user_id: str, draft_id: str, data: dict[str, Any]) -> dict[str, Any]:
now = datetime.now(UTC)
row = RoundtableRecommendHistoryRow(
id=data["id"],
user_id=user_id,
draft_id=draft_id,
objective=(data.get("objective") or "").strip()[:512],
status=(data.get("status") or "done")[:32],
model=(data.get("model") or None),
rationale=data.get("rationale") or None,
picks=_dumps(data.get("picks") or []),
candidates=_dumps(data.get("candidates") or []),
created_at=now,
)
async with self._sf() as session:
session.add(row)
await session.commit()
await session.refresh(row)
return self._history_to_dict(row)