450 lines
17 KiB
Python
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)
|