import asyncio from datetime import datetime, timezone, timedelta from fastapi import APIRouter, Depends, Query from app.internal.auth import require_admin_token from app.internal.feature_flags import get_flags, set_flags from app.internal import firestore as fstore from app.config import settings async def _get_ai_enabled_system_ids(global_flags: dict) -> set[str]: """Return system_ids where at least one AI function (STT or correlation) is effectively on.""" global_stt = global_flags.get("stt_enabled", True) global_corr = global_flags.get("correlation_enabled", True) all_systems = await fstore.collection_list("systems") enabled: set[str] = set() for system in all_systems: sid = system.get("system_id") if not sid: continue ai_flags = system.get("ai_flags") or {} if ai_flags.get("stt_enabled", global_stt) or ai_flags.get("correlation_enabled", global_corr): enabled.add(sid) return enabled router = APIRouter(prefix="/admin", tags=["admin"]) @router.get("/features") async def get_feature_flags(_=Depends(require_admin_token)): """ Return the current AI feature flag state. Admin-only (SAAS_PLAN.md B2c) — was previously any authenticated user via require_firebase_token, which handed platform-wide AI configuration state to every signed-in viewer regardless of org. """ return await get_flags() @router.put("/features") async def update_feature_flags(body: dict, _=Depends(require_admin_token)): """Update one or more AI feature flags. Admin only.""" return await set_flags(body) @router.get("/debug/correlation") async def debug_correlation( limit: int = Query(20, ge=1, le=100), orphan_hours: int = Query(48, ge=1, le=168), ai_systems_only: bool = Query(False, description="Restrict to systems with STT or correlation currently enabled"), _=Depends(require_admin_token), ): """ Return the last N incidents with full correlation debug detail, plus recent orphaned calls. Each incident includes a calls_detail array with per-call corr_* fields so you can see exactly which correlation path fired (or didn't) for every call in the incident. Embeddings are stripped — they're large float arrays and unreadable. Query params: limit — number of incidents to return, sorted by updated_at desc (default 20, max 100) orphan_hours — how far back to scan for orphaned calls (default 48h, max 168h / 1 week) """ def _strip(doc: dict) -> dict: return {k: v for k, v in doc.items() if k != "embedding"} def _call_summary(call: dict) -> dict: return { "call_id": call.get("call_id"), "started_at": call.get("started_at"), "ended_at": call.get("ended_at"), "duration_s": call.get("duration_s"), "talkgroup_id": call.get("talkgroup_id"), "talkgroup_name": call.get("talkgroup_name"), "system_id": call.get("system_id"), "node_id": call.get("node_id"), "incident_type": call.get("incident_type"), "tags": call.get("tags"), "location": call.get("location"), "location_coords": call.get("location_coords"), "units": call.get("units"), "vehicles": call.get("vehicles"), "cleared_units": call.get("cleared_units"), "severity": call.get("severity"), "transcript": call.get("transcript_corrected") or call.get("transcript"), # Correlation decision fields written back by incident_correlator "corr_path": call.get("corr_path"), "corr_incident_idle_min": call.get("corr_incident_idle_min"), "corr_distance_km": call.get("corr_distance_km"), "corr_score": call.get("corr_score"), "corr_candidates": call.get("corr_candidates"), "corr_shared_units": call.get("corr_shared_units"), "corr_fit_signal": call.get("corr_fit_signal"), "corr_matched_units": call.get("corr_matched_units"), "corr_sweep_count": call.get("corr_sweep_count"), "skip_reason": call.get("skip_reason"), # LLM consensus tier fields — written by upload.py's # _correlate_with_consensus / llm_correlator.py, but previously # dropped here, making it impossible to tell from this endpoint # whether the LLM correlation tier is actually running (server-26#24). "corr_consensus": call.get("corr_consensus"), "corr_llm_reasoning": call.get("corr_llm_reasoning"), "corr_llm_action": call.get("corr_llm_action"), "corr_rules_action": call.get("corr_rules_action"), } # ── Determine which systems have AI active ──────────────────────────────── # NOT a filter by default. Restricting to AI-enabled systems meant the view # emptied itself the moment the flags went off — which is precisely when a # window gets reviewed. On 2026-08-23 it dropped from 100 incidents to 6 # between switching correlation off and opening the tab. Pass # ai_systems_only=true to get the old behaviour. global_flags = await get_flags() ai_systems = await _get_ai_enabled_system_ids(global_flags) def _in_scope(system_ids: list) -> bool: if not ai_systems_only: return True return any(sid in ai_systems for sid in system_ids) # ── Fetch recent incidents (AI-enabled systems only) ────────────────────── # Read a bounded, already-sorted window rather than the whole collection. # This route used to pull every incident ever created and sort in Python, # which stopped returning at all once the collection grew — Firestore kills # an unbounded scan with a 503 and the request just hangs. Ordering on the # single field updated_at needs no composite index. # # The AI-system filter runs in Python (it's a membership test against a set # the flags decide), so the window has to be wider than `limit` or filtering # could empty it. 10x with a floor of 200 covers a debug view; if a fetch # still comes back short, incidents_window_exhausted says so in the payload # rather than quietly looking like "no incidents". window = max(limit * 10, 200) all_incidents = await fstore.collection_where( "incidents", [], order_by=[("updated_at", "DESCENDING")], limit_to=window, ) ai_incidents = [i for i in all_incidents if _in_scope(i.get("system_ids") or [])] incidents = ai_incidents[:limit] incidents_window_exhausted = len(all_incidents) >= window and len(ai_incidents) < limit # ── Fetch all linked call docs in parallel ──────────────────────────────── all_call_ids: list[str] = [] for inc in incidents: all_call_ids.extend(inc.get("call_ids") or []) unique_call_ids = list(dict.fromkeys(all_call_ids)) # dedupe, preserve order call_docs = await asyncio.gather(*(fstore.doc_get("calls", cid) for cid in unique_call_ids)) # Key off the id we asked for, not doc["call_id"]. At least one stored call # has no call_id field -- the document id is authoritative and always # present, while the field is written by the upload path and evidently was # not always there. Indexing the field raised KeyError and took the whole # debug view down with a 500 over a single malformed document. call_map: dict[str, dict] = { cid: doc for cid, doc in zip(unique_call_ids, call_docs) if doc } # ── Build incident debug records ────────────────────────────────────────── incident_records = [] for inc in incidents: rec = _strip(inc) rec["calls_detail"] = [ _call_summary(call_map[cid]) for cid in (inc.get("call_ids") or []) if cid in call_map ] incident_records.append(rec) # ── Recent orphaned calls (AI-enabled systems only) ─────────────────────── # Use a single-field range query to avoid requiring a composite Firestore index; # filter status and system in Python. cutoff = datetime.now(timezone.utc) - timedelta(hours=orphan_hours) # Bounded for the same reason as the incident read above. The range and the # sort are both on ended_at, which is what keeps this a single-field query # needing no composite index. _ORPHAN_SCAN_CAP = 3000 recent_calls = await fstore.collection_where( "calls", [("ended_at", ">=", cutoff)], order_by=[("ended_at", "DESCENDING")], limit_to=_ORPHAN_SCAN_CAP, ) orphan_scan_truncated = len(recent_calls) >= _ORPHAN_SCAN_CAP orphans = [ _call_summary(c) for c in recent_calls if c.get("status") == "ended" and not c.get("incident_ids") and not c.get("incident_id") and not c.get("duplicate_of") # another node's copy — never meant to correlate and _in_scope([c.get("system_id")]) ] orphans.sort(key=lambda c: c.get("started_at", ""), reverse=True) # Summarise orphans by talkgroup so the volume and source are immediately visible. orphans_by_tg: dict[str, dict] = {} for o in orphans: tg_key = str(o.get("talkgroup_id") or "unknown") if tg_key not in orphans_by_tg: orphans_by_tg[tg_key] = { "talkgroup_id": o.get("talkgroup_id"), "talkgroup_name": o.get("talkgroup_name") or "unknown", "count": 0, "no_type_count": 0, "sweep_exhausted_count": 0, } orphans_by_tg[tg_key]["count"] += 1 if not o.get("incident_type") and not o.get("tags"): orphans_by_tg[tg_key]["no_type_count"] += 1 if (o.get("corr_sweep_count") or 0) >= 3: orphans_by_tg[tg_key]["sweep_exhausted_count"] += 1 # ── Summary ─────────────────────────────────────────────────────────────── # Everything below was being recomputed by hand from the raw payload on # every review — path counts, how much of the run the LLM tier actually saw, # how many incidents ended up with the "Ems — TGID 9048" fallback name, and # whether anything blew past the server-26#22 caps. Compute it once, here, # where the data already is. def _tally(values) -> dict: out: dict[str, int] = {} for v in values: k = str(v) if v is not None else "none" out[k] = out.get(k, 0) + 1 return dict(sorted(out.items(), key=lambda kv: kv[1], reverse=True)) linked = [c for inc in incident_records for c in (inc.get("calls_detail") or [])] call_counts = [len(inc.get("call_ids") or []) for inc in incident_records] def _span_minutes(inc: dict) -> float: stamps = sorted( s for s in ((c.get("started_at") or "") for c in (inc.get("calls_detail") or [])) if s ) if len(stamps) < 2: return 0.0 try: first = datetime.fromisoformat(str(stamps[0]).replace("Z", "+00:00")) last = datetime.fromisoformat(str(stamps[-1]).replace("Z", "+00:00")) return round((last - first).total_seconds() / 60, 1) except ValueError: return 0.0 spans = [_span_minutes(inc) for inc in incident_records] with_transcript = sum(1 for c in linked if (c.get("transcript") or "").strip()) fallback_titles = sum( 1 for inc in incident_records if " — TGID " in (inc.get("title") or "") or (inc.get("title") or "").endswith("Unknown Talkgroup") ) over_cap = [ {"incident_id": inc.get("incident_id"), "title": inc.get("title"), "calls": len(inc.get("call_ids") or []), "span_minutes": _span_minutes(inc)} for inc in incident_records if len(inc.get("call_ids") or []) > settings.incident_max_calls or _span_minutes(inc) > settings.incident_max_duration_minutes ] summary = { "ai_systems_only": ai_systems_only, "ai_enabled_system_ids": sorted(ai_systems), "linked_call_count": len(linked), "corr_path": _tally(c.get("corr_path") for c in linked), "corr_fit_signal": _tally(c.get("corr_fit_signal") for c in linked), "corr_consensus": _tally(c.get("corr_consensus") for c in linked), "corr_llm_action": _tally(c.get("corr_llm_action") for c in linked), # STT coverage: correlation quality is capped by this, so it belongs in # the same view rather than a separate investigation. "linked_calls_with_transcript": with_transcript, "linked_calls_without_transcript": len(linked) - with_transcript, "orphans_with_transcript": sum(1 for o in orphans if (o.get("transcript") or "").strip()), # Fragmentation vs merging, the two failure directions. "single_call_incidents": sum(1 for n in call_counts if n == 1), "median_calls_per_incident": sorted(call_counts)[len(call_counts) // 2] if call_counts else 0, "max_calls_in_one_incident": max(call_counts) if call_counts else 0, "max_span_minutes": max(spans) if spans else 0.0, "incidents_over_cap": over_cap, "caps": { "incident_max_calls": settings.incident_max_calls, "incident_max_duration_minutes": settings.incident_max_duration_minutes, }, # Titling health — server-26#34. "fallback_titled_incidents": fallback_titles, "titled_incidents": len(incident_records) - fallback_titles, } return { "generated_at": datetime.now(timezone.utc).isoformat(), "summary": summary, # Both reads are capped, so say plainly when a cap was hit — otherwise a # truncated window is indistinguishable from a quiet night. "incidents_window_exhausted": incidents_window_exhausted, "orphan_scan_truncated": orphan_scan_truncated, "orphan_scan_cap": _ORPHAN_SCAN_CAP, "incident_count": len(incident_records), "orphaned_call_count": len(orphans), "orphans_by_talkgroup": sorted(orphans_by_tg.values(), key=lambda x: x["count"], reverse=True), "incidents": incident_records, "orphaned_calls": orphans[:250], } @router.get("/audit") async def get_audit_log( limit: int = Query(50, ge=1, le=200), offset: int = Query(0, ge=0), _=Depends(require_admin_token), ): """Return paginated audit log entries, most recent first.""" entries = await fstore.collection_list("audit_log") entries.sort(key=lambda e: e.get("timestamp", ""), reverse=True) return entries[offset: offset + limit]