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