203 lines
9.0 KiB
Python
203 lines
9.0 KiB
Python
"""SQLAlchemy-backed roundtable business-chain store.
|
||
|
||
The ``seats`` JSON column is stored as a serialized string in
|
||
``PortableLongText`` and round-tripped here with ``json.dumps`` / ``json.loads``
|
||
— same approach as ``roundtable_drafts`` / ``ai_writing_sessions``.
|
||
|
||
All write/read paths are scoped by ``user_id`` for ownership isolation; a row
|
||
that belongs to another user is treated as not found.
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import json
|
||
from datetime import UTC, datetime
|
||
from typing import Any
|
||
|
||
from sqlalchemy import select
|
||
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
|
||
|
||
from deerflow.persistence.roundtable_chains.model import RoundtableChainRow
|
||
|
||
_UNSET = object()
|
||
|
||
|
||
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 RoundtableChainRepository:
|
||
def __init__(self, session_factory: async_sessionmaker[AsyncSession]) -> None:
|
||
self._sf = session_factory
|
||
|
||
@staticmethod
|
||
def _to_dict(row: RoundtableChainRow, current_user_id: str | None = None) -> dict[str, Any]:
|
||
return {
|
||
"id": row.id,
|
||
"title": row.title,
|
||
"description": row.description,
|
||
"seats": _loads(row.seats) or [],
|
||
# None when absent → frontend reads it as "linear chain" (no layering).
|
||
"stages": _loads(row.stages),
|
||
# Per-stage goals aligned to ``stages``; None when absent.
|
||
"stage_goals": _loads(getattr(row, "stage_goals", None)),
|
||
# 总控编排提示(写给总控的派活指引,dag-only);None/空 = 不注入。纯文本,不 json。
|
||
"coordinator_prompt": getattr(row, "coordinator_prompt", None),
|
||
# 默认席位执行模式(flash/thinking/pro/ultra);None = 未设。
|
||
"seat_mode": getattr(row, "seat_mode", None),
|
||
# 研讨前并行取数开关(逐条 opt-in);缺省 False。
|
||
"gather_first": bool(getattr(row, "gather_first", False)),
|
||
# 发布状态 + 是否本人所有(前端据此显示「公共」徽标、决定能否编辑 / 发布)。
|
||
"is_public": bool(getattr(row, "is_public", False)),
|
||
# 启停状态:缺省 True(存量行 / 补列前视为启用)。
|
||
"enabled": bool(getattr(row, "enabled", True)),
|
||
# 「待审核」标记:True = owner 已提交审核、等待管理员发布;发布时清回 False。
|
||
"pending_review": bool(getattr(row, "pending_review", False)),
|
||
"is_mine": (current_user_id is None) or (row.user_id == current_user_id),
|
||
# 配置者(owner)的 user_id;路由层据此解析成可读用户名(owner_name)。
|
||
"owner_id": row.user_id,
|
||
"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,
|
||
}
|
||
|
||
async def list_chains(
|
||
self, user_id: str, *, scope: str = "mine", limit: int = 200, is_admin: bool = False
|
||
) -> list[dict[str, Any]]:
|
||
"""列出链条。
|
||
|
||
``scope``:
|
||
- ``"mine"``(默认):仅当前用户自己的链条;
|
||
- ``"public"``:所有**已发布**(``is_public=True``)链条,含本人发布的;
|
||
- ``"pending"``:所有**待审核**(``pending_review=True``)链条,**仅管理员**(由路由层
|
||
校验后传 ``is_admin=True``;非管理员误传则按 ``"mine"`` 退化,绝不越权);
|
||
- ``"all"``:所有账号的全部链条(**仅管理员**,由路由层校验后传 ``is_admin=True``;
|
||
非管理员误传则按 ``"mine"`` 退化,绝不越权)。
|
||
其余值按 ``"mine"`` 处理。``is_mine`` 标记每条是否本人所有,供前端判断可否编辑。
|
||
"""
|
||
if scope == "public":
|
||
cond = RoundtableChainRow.is_public.is_(True)
|
||
elif scope == "pending" and is_admin:
|
||
cond = RoundtableChainRow.pending_review.is_(True)
|
||
elif scope == "all" and is_admin:
|
||
cond = None
|
||
else:
|
||
cond = RoundtableChainRow.user_id == user_id
|
||
stmt = select(RoundtableChainRow)
|
||
if cond is not None:
|
||
stmt = stmt.where(cond)
|
||
stmt = stmt.order_by(RoundtableChainRow.updated_at.desc()).limit(limit)
|
||
async with self._sf() as session:
|
||
result = await session.execute(stmt)
|
||
return [self._to_dict(row, user_id) for row in result.scalars()]
|
||
|
||
async def get_chain(
|
||
self, chain_id: str, user_id: str, *, is_admin: bool = False
|
||
) -> dict[str, Any] | None:
|
||
"""取一条链条:owner 可见,**已发布的公共链条任何人可见**(供选入研讨);
|
||
**管理员可见任意链条**(``is_admin=True``,供后台管理配置)。"""
|
||
async with self._sf() as session:
|
||
row = await session.get(RoundtableChainRow, chain_id)
|
||
if row is None:
|
||
return None
|
||
if not is_admin and row.user_id != user_id and not bool(getattr(row, "is_public", False)):
|
||
return None
|
||
return self._to_dict(row, user_id)
|
||
|
||
async def create_chain(self, user_id: str, data: dict[str, Any]) -> dict[str, Any]:
|
||
now = datetime.now(UTC)
|
||
row = RoundtableChainRow(
|
||
id=data["id"],
|
||
user_id=user_id,
|
||
title=(data.get("title") or "").strip(),
|
||
description=(data.get("description") or None),
|
||
seats=_dumps(data.get("seats") or []),
|
||
stages=_dumps(data.get("stages")),
|
||
stage_goals=_dumps(data.get("stage_goals")),
|
||
coordinator_prompt=(data.get("coordinator_prompt") or None),
|
||
seat_mode=(data.get("seat_mode") or None),
|
||
gather_first=bool(data.get("gather_first") or False),
|
||
is_public=bool(data.get("is_public") or False),
|
||
enabled=True if data.get("enabled") is None else bool(data.get("enabled")),
|
||
pending_review=bool(data.get("pending_review") or False),
|
||
created_at=now,
|
||
updated_at=now,
|
||
)
|
||
async with self._sf() as session:
|
||
session.add(row)
|
||
await session.commit()
|
||
await session.refresh(row)
|
||
return self._to_dict(row, user_id)
|
||
|
||
async def update_chain(
|
||
self,
|
||
chain_id: str,
|
||
user_id: str,
|
||
*,
|
||
title=_UNSET,
|
||
description=_UNSET,
|
||
seats=_UNSET,
|
||
stages=_UNSET,
|
||
stage_goals=_UNSET,
|
||
coordinator_prompt=_UNSET,
|
||
seat_mode=_UNSET,
|
||
gather_first=_UNSET,
|
||
is_public=_UNSET,
|
||
enabled=_UNSET,
|
||
pending_review=_UNSET,
|
||
is_admin: bool = False,
|
||
) -> dict[str, Any] | None:
|
||
async with self._sf() as session:
|
||
row = await session.get(RoundtableChainRow, chain_id)
|
||
# 管理员可改任意链条;否则仅 owner 本人。不变更 user_id(不转移归属)。
|
||
if row is None or (not is_admin and row.user_id != user_id):
|
||
return None
|
||
if title is not _UNSET and title is not None:
|
||
row.title = str(title).strip()
|
||
if description is not _UNSET:
|
||
row.description = description or None
|
||
if seats is not _UNSET:
|
||
row.seats = _dumps(seats)
|
||
if stages is not _UNSET:
|
||
row.stages = _dumps(stages)
|
||
if stage_goals is not _UNSET:
|
||
row.stage_goals = _dumps(stage_goals)
|
||
if coordinator_prompt is not _UNSET:
|
||
row.coordinator_prompt = coordinator_prompt or None
|
||
if seat_mode is not _UNSET:
|
||
row.seat_mode = seat_mode or None
|
||
if gather_first is not _UNSET and gather_first is not None:
|
||
row.gather_first = bool(gather_first)
|
||
if is_public is not _UNSET and is_public is not None:
|
||
row.is_public = bool(is_public)
|
||
if enabled is not _UNSET and enabled is not None:
|
||
row.enabled = bool(enabled)
|
||
if pending_review is not _UNSET and pending_review is not None:
|
||
row.pending_review = bool(pending_review)
|
||
row.updated_at = datetime.now(UTC)
|
||
await session.commit()
|
||
await session.refresh(row)
|
||
return self._to_dict(row, user_id)
|
||
|
||
async def delete_chain(self, chain_id: str, user_id: str, *, is_admin: bool = False) -> bool:
|
||
async with self._sf() as session:
|
||
row = await session.get(RoundtableChainRow, chain_id)
|
||
# 管理员可删任意链条;否则仅 owner 本人。
|
||
if row is None or (not is_admin and row.user_id != user_id):
|
||
return False
|
||
await session.delete(row)
|
||
await session.commit()
|
||
return True
|