"""Admin-only live concurrency monitor for in-flight model runs. Surfaces every *currently running* call into the agent runtime (interactive chat, agent chat, roundtable seats, …) so an administrator can see how much concurrent pressure the system is putting on the LLM and **truly stop** any of them — ``RunManager.cancel`` aborts the run's asyncio task, which cancels the in-flight model stream, immediately freeing that capacity. All endpoints require ``system_role == "admin"`` (shared ``_require_admin``). The data source is the in-memory :class:`RunManager` registry (``list_active``), enriched best-effort with the owning user's email + the agent's display name from the DB. Enrichment failures degrade gracefully (the run still shows, just with rawer labels) — the monitor must never break because a lookup failed. """ from __future__ import annotations import logging from datetime import UTC, datetime, timedelta from fastapi import APIRouter, Query, Request from pydantic import BaseModel from sqlalchemy import select from app.gateway.deps import get_concurrency_gate, get_concurrency_sample_store, get_run_manager from app.gateway.routers.admin_users import _require_admin from deerflow.config.system_settings import load_system_settings, save_system_settings from deerflow.persistence.agents.model import AgentRow from deerflow.persistence.engine import get_session_factory from deerflow.persistence.llm_metrics.model import LlmCallMetricRow from deerflow.persistence.thread_meta.model import ThreadMetaRow from deerflow.persistence.user.model import UserRow logger = logging.getLogger(__name__) router = APIRouter(prefix="/api/admin/active-runs", tags=["admin-active-runs"]) # Lookback windows are clamped so an aggregation never scans an unbounded range. _MAX_LOOKBACK_HOURS = 24 * 14 # 14 days # Per-granularity default lookback (hours): minute view = recent, hour view = wider. _DEFAULT_LOOKBACK_HOURS = {"minute": 3, "hour": 48} # assistant_id values that are not real custom-agent ids (the default lead agent). _DEFAULT_ASSISTANT_IDS = {"", "default", "lead_agent", "agent"} class ActiveRunItem(BaseModel): run_id: str thread_id: str status: str created_at: str = "" elapsed_seconds: int = 0 user_id: str = "" user_name: str = "" agent_id: str = "" agent_name: str = "" model_name: str = "" thread_title: str = "" kind: str = "对话" multitask_strategy: str = "reject" class ActiveRunsResponse(BaseModel): count: int server_time: str runs: list[ActiveRunItem] class CancelRunResponse(BaseModel): run_id: str cancelled: bool class CancelAllResponse(BaseModel): cancelled_count: int class ConcurrencyPoint(BaseModel): bucket: str peak: int avg: float samples: int class ConcurrencySeriesResponse(BaseModel): granularity: str since: str until: str enabled: bool points: list[ConcurrencyPoint] class TokenSpeedPoint(BaseModel): bucket: str avg_tokens_per_sec: float min_tokens_per_sec: float max_tokens_per_sec: float samples: int class TokenSpeedSeriesResponse(BaseModel): granularity: str since: str until: str points: list[TokenSpeedPoint] class MonitorSettingsResponse(BaseModel): enabled: bool def _resolve_window(granularity: str, hours: float | None) -> tuple[str, datetime, datetime]: """Clamp the granularity + lookback window into (gran, since, until).""" gran = "hour" if granularity == "hour" else "minute" lookback = hours if hours and hours > 0 else _DEFAULT_LOOKBACK_HOURS[gran] lookback = min(float(lookback), _MAX_LOOKBACK_HOURS) until = datetime.now(UTC) since = until - timedelta(hours=lookback) return gran, since, until def _bucket_start(dt: datetime, granularity: str) -> datetime: if granularity == "hour": return dt.replace(minute=0, second=0, microsecond=0) return dt.replace(second=0, microsecond=0) def _derive_tokens_per_sec(tps, output_tokens, total_tokens, duration_ms) -> float | None: """Use the stored tokens/sec, else derive from tokens ÷ duration.""" if tps not in (None, 0): return float(tps) numerator = output_tokens if output_tokens else total_tokens if not numerator or not duration_ms or duration_ms <= 0: return None return round((float(numerator) / float(duration_ms)) * 1000.0, 3) def _elapsed_seconds(created_at: str, now: datetime) -> int: """Seconds since the (UTC ISO) ``created_at`` stamp; 0 on any parse error.""" if not created_at: return 0 try: started = datetime.fromisoformat(created_at) except (ValueError, TypeError): return 0 if started.tzinfo is None: started = started.replace(tzinfo=UTC) delta = (now - started).total_seconds() return max(0, int(delta)) def _elapsed_seconds_from_epoch(created_at: float, now: datetime) -> int: try: if created_at <= 0: return 0 return max(0, int(now.timestamp() - float(created_at))) except Exception: return 0 def _kind_for(metadata: dict | None) -> str: """Coarse human label for the run's origin, from its thread metadata.""" meta = metadata or {} thread_type = str(meta.get("thread_type") or "") if thread_type == "roundtable": return "圆桌会商" if meta.get("system"): return "系统/后台" if meta.get("taskId"): return "任务对话" return "对话" async def _enrich(records) -> dict: """Batch-resolve thread metadata, user emails and agent names. Returns a dict of lookup maps; every query is wrapped so a failure just yields empty maps (the runs still render with raw ids). """ thread_ids = {r.thread_id for r in records if r.thread_id} assistant_ids = { r.assistant_id for r in records if r.assistant_id and r.assistant_id not in _DEFAULT_ASSISTANT_IDS } thread_meta: dict[str, ThreadMetaRow] = {} user_emails: dict[str, str] = {} agent_names: dict[str, str] = {} try: session_factory = get_session_factory() async with session_factory() as session: if thread_ids: rows = ( await session.execute( select(ThreadMetaRow).where(ThreadMetaRow.thread_id.in_(thread_ids)) ) ).scalars().all() thread_meta = {row.thread_id: row for row in rows} user_ids = {row.user_id for row in thread_meta.values() if row.user_id} if user_ids: rows = ( await session.execute( select(UserRow.id, UserRow.email).where(UserRow.id.in_(user_ids)) ) ).all() user_emails = {r[0]: (r[1] or "") for r in rows} if assistant_ids: rows = ( await session.execute( select(AgentRow.id, AgentRow.name).where(AgentRow.id.in_(assistant_ids)) ) ).all() agent_names = {r[0]: (r[1] or "") for r in rows} except Exception: # pragma: no cover - enrichment is best-effort logger.warning("Failed to enrich active runs (non-fatal)", exc_info=True) return {"thread_meta": thread_meta, "user_emails": user_emails, "agent_names": agent_names} @router.get("", response_model=ActiveRunsResponse) @router.get("/", response_model=ActiveRunsResponse) async def list_active_runs(request: Request) -> ActiveRunsResponse: """List every in-flight run with who/what is driving it (admin only).""" await _require_admin(request) run_mgr = get_run_manager(request) gate = get_concurrency_gate(request) leases = await gate.list_active() if gate is not None else None records = [] if leases is not None else await run_mgr.list_active() now = datetime.now(UTC) enrich_source = leases if leases is not None else records maps = await _enrich(enrich_source) thread_meta = maps["thread_meta"] user_emails = maps["user_emails"] agent_names = maps["agent_names"] items: list[ActiveRunItem] = [] if leases is not None: for r in leases: meta_row = thread_meta.get(r.thread_id) user_id = r.user_id or (meta_row.user_id if meta_row else "") or "" metadata = (meta_row.metadata_json if meta_row else None) or {} assistant_id = r.assistant_id or "" if assistant_id in _DEFAULT_ASSISTANT_IDS: agent_name = "榛樿鍔╂墜" else: agent_name = agent_names.get(assistant_id) or assistant_id items.append( ActiveRunItem( run_id=r.run_id, thread_id=r.thread_id, status="running", created_at=datetime.fromtimestamp(r.created_at, tz=UTC).isoformat() if r.created_at else "", elapsed_seconds=_elapsed_seconds_from_epoch(r.created_at, now), user_id=user_id, user_name=user_emails.get(user_id) or user_id or "鏈煡鐢ㄦ埛", agent_id=assistant_id, agent_name=agent_name, model_name=r.model_name, thread_title=(meta_row.display_name if meta_row else "") or "", kind=_kind_for(metadata), multitask_strategy="reject", ) ) return ActiveRunsResponse(count=len(items), server_time=now.isoformat(), runs=items) for r in records: meta_row = thread_meta.get(r.thread_id) user_id = (meta_row.user_id if meta_row else "") or "" metadata = (meta_row.metadata_json if meta_row else None) or r.metadata or {} assistant_id = r.assistant_id or "" if assistant_id in _DEFAULT_ASSISTANT_IDS: agent_name = "默认助手" else: agent_name = agent_names.get(assistant_id) or assistant_id items.append( ActiveRunItem( run_id=r.run_id, thread_id=r.thread_id, status=r.status.value, created_at=r.created_at, elapsed_seconds=_elapsed_seconds(r.created_at, now), user_id=user_id, user_name=user_emails.get(user_id) or user_id or "未知用户", agent_id=assistant_id, agent_name=agent_name, model_name=str((r.metadata or {}).get("model_name") or ""), thread_title=(meta_row.display_name if meta_row else "") or "", kind=_kind_for(metadata), multitask_strategy=r.multitask_strategy, ) ) return ActiveRunsResponse(count=len(items), server_time=now.isoformat(), runs=items) @router.post("/{run_id}/cancel", response_model=CancelRunResponse) async def cancel_active_run( run_id: str, request: Request, action: str = Query(default="interrupt", description="interrupt 保留检查点 / rollback 回滚"), ) -> CancelRunResponse: """Truly stop one in-flight run — cancels its task and the live model stream.""" await _require_admin(request) run_mgr = get_run_manager(request) normalized = "rollback" if action == "rollback" else "interrupt" cancelled = await run_mgr.cancel(run_id, action=normalized) return CancelRunResponse(run_id=run_id, cancelled=cancelled) @router.post("/cancel-all", response_model=CancelAllResponse) async def cancel_all_active_runs(request: Request) -> CancelAllResponse: """Stop every in-flight run at once (admin emergency relief valve).""" await _require_admin(request) run_mgr = get_run_manager(request) records = await run_mgr.list_active() cancelled = 0 for r in records: if await run_mgr.cancel(r.run_id, action="interrupt"): cancelled += 1 return CancelAllResponse(cancelled_count=cancelled) @router.get("/concurrency", response_model=ConcurrencySeriesResponse) async def concurrency_series( request: Request, granularity: str = Query(default="minute", description="minute | hour"), hours: float | None = Query(default=None, description="回看小时数(默认按粒度取)"), ) -> ConcurrencySeriesResponse: """Per-minute / per-hour 并发量 series (peak + avg) for the pressure chart.""" await _require_admin(request) gran, since, until = _resolve_window(granularity, hours) enabled = load_system_settings().concurrency_monitor.enabled store = get_concurrency_sample_store(request) points: list[ConcurrencyPoint] = [] if store is not None: buckets = await store.aggregate(since=since, until=until, granularity=gran) points = [ ConcurrencyPoint(bucket=b.bucket, peak=b.peak, avg=b.avg, samples=b.samples) for b in buckets ] return ConcurrencySeriesResponse( granularity=gran, since=since.isoformat(), until=until.isoformat(), enabled=enabled, points=points, ) @router.get("/token-speed", response_model=TokenSpeedSeriesResponse) async def token_speed_series( request: Request, granularity: str = Query(default="minute", description="minute | hour"), hours: float | None = Query(default=None, description="回看小时数(默认按粒度取)"), ) -> TokenSpeedSeriesResponse: """Per-minute / per-hour 大模型出字速度 (tokens/sec) series. Aggregated from the persisted LLM-call metrics: a dip in tokens/sec for a bucket points at the model/provider being the bottleneck for that window. """ await _require_admin(request) gran, since, until = _resolve_window(granularity, hours) buckets: dict[datetime, dict[str, float]] = {} try: session_factory = get_session_factory() if session_factory is not None: async with session_factory() as session: rows = ( await session.execute( select( LlmCallMetricRow.created_at, LlmCallMetricRow.tokens_per_sec, LlmCallMetricRow.output_tokens, LlmCallMetricRow.total_tokens, LlmCallMetricRow.duration_ms, ) .where( LlmCallMetricRow.created_at >= since, LlmCallMetricRow.created_at <= until, LlmCallMetricRow.status == "success", ) .order_by(LlmCallMetricRow.created_at) .limit(100_000) ) ).all() for created_at, tps_raw, output_tokens, total_tokens, duration_ms in rows: tps = _derive_tokens_per_sec(tps_raw, output_tokens, total_tokens, duration_ms) if tps is None or tps <= 0 or created_at is None: continue if created_at.tzinfo is None: created_at = created_at.replace(tzinfo=UTC) key = _bucket_start(created_at, gran) entry = buckets.get(key) if entry is None: buckets[key] = {"sum": tps, "min": tps, "max": tps, "count": 1} else: entry["sum"] += tps entry["min"] = min(entry["min"], tps) entry["max"] = max(entry["max"], tps) entry["count"] += 1 except Exception: # pragma: no cover - chart must degrade gracefully logger.warning("Failed to aggregate token-speed series (non-fatal)", exc_info=True) points: list[TokenSpeedPoint] = [] for key in sorted(buckets): entry = buckets[key] count = max(1, int(entry["count"])) points.append( TokenSpeedPoint( bucket=key.isoformat(), avg_tokens_per_sec=round(entry["sum"] / count, 2), min_tokens_per_sec=round(entry["min"], 2), max_tokens_per_sec=round(entry["max"], 2), samples=count, ) ) return TokenSpeedSeriesResponse( granularity=gran, since=since.isoformat(), until=until.isoformat(), points=points, ) @router.get("/settings", response_model=MonitorSettingsResponse) async def get_monitor_settings(request: Request) -> MonitorSettingsResponse: """Read the concurrency-monitor on/off switch (admin only).""" await _require_admin(request) return MonitorSettingsResponse(enabled=load_system_settings().concurrency_monitor.enabled) @router.put("/settings", response_model=MonitorSettingsResponse) async def update_monitor_settings( request: Request, body: MonitorSettingsResponse, ) -> MonitorSettingsResponse: """Toggle the background concurrency sampler on/off (admin only, no restart).""" await _require_admin(request) settings = load_system_settings() updated = settings.model_copy( update={"concurrency_monitor": settings.concurrency_monitor.model_copy(update={"enabled": bool(body.enabled)})} ) save_system_settings(updated) return MonitorSettingsResponse(enabled=updated.concurrency_monitor.enabled)