105 lines
3.9 KiB
Python
105 lines
3.9 KiB
Python
"""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)
|