285 lines
12 KiB
Python
285 lines
12 KiB
Python
"""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)
|