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

203 lines
9.0 KiB
Python
Raw Permalink 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 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