deerflow-code/offline-backend-20260512/backend/app/gateway/routers/admin_active_runs.py
2026-09-07 18:24:55 +08:00

450 lines
17 KiB
Python

"""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)