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

105 lines
3.9 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 大屏绘制智能体 session store (按 taskId,不分权).
``transcript`` / ``report_json`` are JSON serialized into ``PortableLongText``
and round-tripped here with ``json.dumps`` / ``json.loads`` — same approach as
``roundtable_drafts``. All access keys on ``task_id`` and ignores ``user_id``
(open read + shared writes); ``user_id`` is recorded for audit only.
"""
from __future__ import annotations
import json
from datetime import UTC, datetime
from typing import Any
from uuid import uuid4
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
from deerflow.persistence.dashboard_sessions.model import DashboardSessionRow
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
if isinstance(value, str):
return value # already a JSON string (report_json) — store verbatim
return json.dumps(value, ensure_ascii=False)
class DashboardSessionRepository:
def __init__(self, session_factory: async_sessionmaker[AsyncSession]) -> None:
self._sf = session_factory
@staticmethod
def _to_dict(row: DashboardSessionRow) -> dict[str, Any]:
return {
"id": row.id,
"task_id": row.task_id,
"user_id": row.user_id,
"transcript": _loads(row.transcript),
# report_json is a structured-data JSON string; return it as a string
# so the frontend parses it with the same robust parser it uses live.
"report_json": row.report_json,
"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 get_by_task(self, task_id: str) -> dict[str, Any] | None:
"""取该 task 的会话(**忽略 user_id**,不分权)。无则 None。"""
stmt = select(DashboardSessionRow).where(DashboardSessionRow.task_id == task_id).limit(1)
async with self._sf() as session:
result = await session.execute(stmt)
row = result.scalars().first()
return self._to_dict(row) if row is not None else None
async def upsert_by_task(
self,
task_id: str,
*,
transcript: Any = None,
report_json: Any = None,
user_id: str | None = None,
) -> dict[str, Any]:
"""按 task_id upsert:无则建、有则更新(**不分权**)。
``transcript`` / ``report_json`` 仅在传入非 None 时覆盖,便于只更新其中一项。
"""
now = datetime.now(UTC)
async with self._sf() as session:
stmt = select(DashboardSessionRow).where(DashboardSessionRow.task_id == task_id).limit(1)
row = (await session.execute(stmt)).scalars().first()
if row is None:
row = DashboardSessionRow(
id=uuid4().hex,
task_id=task_id,
user_id=user_id,
transcript=_dumps(transcript),
report_json=_dumps(report_json),
created_at=now,
updated_at=now,
)
session.add(row)
else:
if transcript is not None:
row.transcript = _dumps(transcript)
if report_json is not None:
row.report_json = _dumps(report_json)
if user_id is not None:
row.user_id = user_id
row.updated_at = now
await session.commit()
await session.refresh(row)
return self._to_dict(row)