"""Tests for the in-memory MCP per-tool statistics tracker. These pin the two properties the status panel depends on: counters stay exact under the concurrent thread dispatch the mcp SDK uses, and recording never raises into a live tool call. """ from __future__ import annotations import threading from app.mcp.stats import MCPStatsTracker, MCPToolStats, get_mcp_stats_tracker def test_avg_duration_is_none_before_any_call(): """A never-called tool reports None, not 0.0, so the UI can distinguish them.""" assert MCPToolStats().avg_duration_ms is None def test_avg_duration_is_the_mean_of_recorded_durations(): """Average is computed over all calls, successful or not.""" tracker = MCPStatsTracker() for duration in (100.0, 200.0, 300.0): tracker.record(tool="search_regulations", duration_ms=duration, success=True) stats = tracker.snapshot()["search_regulations"] assert stats.calls == 3 assert stats.avg_duration_ms == 200.0 def test_failures_increment_both_calls_and_errors(): """errors is a subset of calls, so the UI can show "2 of 3 failed" honestly.""" tracker = MCPStatsTracker() tracker.record(tool="t", duration_ms=1.0, success=True) tracker.record(tool="t", duration_ms=1.0, success=False) tracker.record(tool="t", duration_ms=1.0, success=False) stats = tracker.snapshot()["t"] assert stats.calls == 3 assert stats.errors == 2 def test_last_called_at_is_set_and_timezone_aware(): """The panel renders this as a local time, which requires an aware datetime.""" tracker = MCPStatsTracker() tracker.record(tool="t", duration_ms=1.0, success=True) last_called = tracker.snapshot()["t"].last_called_at assert last_called is not None assert last_called.tzinfo is not None def test_concurrent_record_calls_are_not_lost(): """8 threads x 100 calls must total exactly 800. Without the lock this loses increments non-deterministically. The mcp SDK runs synchronous tool bodies via anyio.to_thread.run_sync, so this is the real dispatch model, not a hypothetical. """ tracker = MCPStatsTracker() def hammer() -> None: for _ in range(100): tracker.record(tool="search_regulations", duration_ms=1.0, success=True) threads = [threading.Thread(target=hammer) for _ in range(8)] for thread in threads: thread.start() for thread in threads: thread.join() stats = tracker.snapshot()["search_regulations"] assert stats.calls == 800 assert stats.total_duration_ms == 800.0 def test_record_swallows_bad_input_instead_of_raising(): """A malformed duration must not propagate into the caller's tool call.""" tracker = MCPStatsTracker() tracker.record(tool="t", duration_ms="not-a-number", success=True) # type: ignore[arg-type] # Coercion happens before the lock is taken, so the entry is never created # in a half-updated state. assert tracker.snapshot() == {} def test_snapshot_is_a_copy_not_the_live_dict(): """Callers mutating the snapshot must not corrupt the tracker.""" tracker = MCPStatsTracker() tracker.record(tool="t", duration_ms=1.0, success=True) snapshot = tracker.snapshot() snapshot.clear() assert "t" in tracker.snapshot() def test_get_mcp_stats_tracker_returns_a_singleton(): """Instrumentation and the status route must observe the same counters.""" assert get_mcp_stats_tracker() is get_mcp_stats_tracker()