212 lines
7.4 KiB
Python
212 lines
7.4 KiB
Python
"""SQLAlchemy-backed 圆桌会商诊断日志 store。
|
||
|
||
写入(``record``)best-effort、全局共享;读取(``list_logs`` / ``count`` / ``facets``)
|
||
带多维筛选,供管理员诊断页用。``detail`` 以 JSON 序列化进 ``PortableLongText``。
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import json
|
||
from datetime import datetime
|
||
from typing import Any
|
||
from uuid import uuid4
|
||
|
||
from sqlalchemy import delete, func, select
|
||
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
|
||
|
||
from deerflow.persistence.roundtable_diagnostics.model import RoundtableDiagnosticRow
|
||
|
||
|
||
def _dumps(value: Any) -> str | None:
|
||
if value is None:
|
||
return None
|
||
if isinstance(value, str):
|
||
return value
|
||
try:
|
||
return json.dumps(value, ensure_ascii=False, default=str)
|
||
except Exception:
|
||
return str(value)
|
||
|
||
|
||
def _loads(value: Any) -> Any:
|
||
if not isinstance(value, str) or not value:
|
||
return value or None
|
||
try:
|
||
return json.loads(value)
|
||
except Exception:
|
||
return value # 非 JSON 文本(如纯字符串 detail)原样返回
|
||
|
||
|
||
class RoundtableDiagnosticRepository:
|
||
def __init__(self, session_factory: async_sessionmaker[AsyncSession]) -> None:
|
||
self._sf = session_factory
|
||
|
||
@staticmethod
|
||
def _to_dict(row: RoundtableDiagnosticRow) -> dict[str, Any]:
|
||
return {
|
||
"id": row.id,
|
||
"scope": row.scope,
|
||
"stage": row.stage,
|
||
"level": row.level,
|
||
"event": row.event,
|
||
"message": row.message,
|
||
"detail": _loads(row.detail),
|
||
"job_id": row.job_id,
|
||
"draft_id": row.draft_id,
|
||
"task_id": row.task_id,
|
||
"user_id": row.user_id,
|
||
"agent_id": row.agent_id,
|
||
"agent_name": row.agent_name,
|
||
"cycle": row.cycle,
|
||
"created_at": row.created_at.isoformat() if isinstance(row.created_at, datetime) else row.created_at,
|
||
}
|
||
|
||
async def record(
|
||
self,
|
||
*,
|
||
scope: str = "background",
|
||
stage: str = "job",
|
||
level: str = "info",
|
||
event: str = "",
|
||
message: str | None = None,
|
||
detail: Any = None,
|
||
job_id: str | None = None,
|
||
draft_id: str | None = None,
|
||
task_id: str | None = None,
|
||
user_id: str | None = None,
|
||
agent_id: str | None = None,
|
||
agent_name: str | None = None,
|
||
cycle: int | None = None,
|
||
) -> dict[str, Any]:
|
||
"""插入一条诊断事件,返回序列化 dict。"""
|
||
row = RoundtableDiagnosticRow(
|
||
id=uuid4().hex,
|
||
scope=(scope or "background")[:16],
|
||
stage=(stage or "job")[:32],
|
||
level=(level or "info")[:16],
|
||
event=(event or "")[:64],
|
||
message=message,
|
||
detail=_dumps(detail),
|
||
job_id=job_id,
|
||
draft_id=draft_id,
|
||
task_id=task_id,
|
||
user_id=user_id,
|
||
agent_id=(agent_id or None) if agent_id is None else str(agent_id)[:128],
|
||
agent_name=(agent_name or None) if agent_name is None else str(agent_name)[:255],
|
||
cycle=cycle,
|
||
)
|
||
async with self._sf() as session:
|
||
session.add(row)
|
||
await session.commit()
|
||
await session.refresh(row)
|
||
return self._to_dict(row)
|
||
|
||
def _apply_filters(
|
||
self,
|
||
stmt,
|
||
*,
|
||
scope: str | list[str] | None,
|
||
stage: str | None,
|
||
level: str | None,
|
||
event: str | None,
|
||
job_id: str | None,
|
||
draft_id: str | None,
|
||
task_id: str | None,
|
||
user_id: str | None,
|
||
q: str | None,
|
||
since: datetime | None,
|
||
until: datetime | None,
|
||
):
|
||
R = RoundtableDiagnosticRow
|
||
if scope:
|
||
if isinstance(scope, (list, tuple, set)):
|
||
vals = [s for s in scope if s]
|
||
if vals:
|
||
stmt = stmt.where(R.scope.in_(vals))
|
||
else:
|
||
stmt = stmt.where(R.scope == scope)
|
||
if stage:
|
||
stmt = stmt.where(R.stage == stage)
|
||
if level:
|
||
stmt = stmt.where(R.level == level)
|
||
if event:
|
||
stmt = stmt.where(R.event == event)
|
||
if job_id:
|
||
stmt = stmt.where(R.job_id == job_id)
|
||
if draft_id:
|
||
stmt = stmt.where(R.draft_id == draft_id)
|
||
if task_id:
|
||
stmt = stmt.where(R.task_id == task_id)
|
||
if user_id:
|
||
stmt = stmt.where(R.user_id == user_id)
|
||
if q:
|
||
like = f"%{q}%"
|
||
stmt = stmt.where(R.message.ilike(like))
|
||
if since is not None:
|
||
stmt = stmt.where(R.created_at >= since)
|
||
if until is not None:
|
||
stmt = stmt.where(R.created_at <= until)
|
||
return stmt
|
||
|
||
async def list_logs(
|
||
self,
|
||
*,
|
||
scope: str | list[str] | None = None,
|
||
stage: str | None = None,
|
||
level: str | None = None,
|
||
event: str | None = None,
|
||
job_id: str | None = None,
|
||
draft_id: str | None = None,
|
||
task_id: str | None = None,
|
||
user_id: str | None = None,
|
||
q: str | None = None,
|
||
since: datetime | None = None,
|
||
until: datetime | None = None,
|
||
limit: int = 100,
|
||
offset: int = 0,
|
||
) -> tuple[list[dict[str, Any]], int]:
|
||
"""筛选 + 分页(newest-first)。返回 ``(items, total)``。"""
|
||
R = RoundtableDiagnosticRow
|
||
base = self._apply_filters(
|
||
select(R),
|
||
scope=scope, stage=stage, level=level, event=event,
|
||
job_id=job_id, draft_id=draft_id, task_id=task_id, user_id=user_id,
|
||
q=q, since=since, until=until,
|
||
)
|
||
count_stmt = self._apply_filters(
|
||
select(func.count()).select_from(R),
|
||
scope=scope, stage=stage, level=level, event=event,
|
||
job_id=job_id, draft_id=draft_id, task_id=task_id, user_id=user_id,
|
||
q=q, since=since, until=until,
|
||
)
|
||
page_stmt = base.order_by(R.created_at.desc(), R.id.desc()).limit(limit).offset(offset)
|
||
async with self._sf() as session:
|
||
total = int((await session.execute(count_stmt)).scalar() or 0)
|
||
rows = (await session.execute(page_stmt)).scalars().all()
|
||
return [self._to_dict(r) for r in rows], total
|
||
|
||
async def facets(self) -> dict[str, list[str]]:
|
||
"""返回各筛选维度的 distinct 取值(供前端下拉),供筛选 UI 用。"""
|
||
R = RoundtableDiagnosticRow
|
||
out: dict[str, list[str]] = {}
|
||
async with self._sf() as session:
|
||
for key, col in (("scope", R.scope), ("stage", R.stage), ("level", R.level), ("event", R.event)):
|
||
rows = (await session.execute(select(col).distinct())).scalars().all()
|
||
out[key] = sorted([str(v) for v in rows if v])
|
||
return out
|
||
|
||
async def list_older_than(self, cutoff: datetime, *, limit: int = 1000) -> list[str]:
|
||
R = RoundtableDiagnosticRow
|
||
stmt = select(R.id).where(R.created_at < cutoff).limit(limit)
|
||
async with self._sf() as session:
|
||
return [str(v) for v in (await session.execute(stmt)).scalars().all()]
|
||
|
||
async def delete_by_ids(self, ids: list[str]) -> int:
|
||
if not ids:
|
||
return 0
|
||
R = RoundtableDiagnosticRow
|
||
async with self._sf() as session:
|
||
result = await session.execute(delete(R).where(R.id.in_(ids)))
|
||
await session.commit()
|
||
return int(result.rowcount or 0)
|