feat: add PostgresModelUsageStore and ModelUsageTracker.seed()
- Add seed() method to ModelUsageTracker for bulk-loading persisted entries at startup - Create PostgresModelUsageStore for persistence of model usage counters to Postgres - Store only current cumulative snapshots (no historical time-series) - Use standard CREATE TABLE IF NOT EXISTS idiom matching other Postgres stores - Add comprehensive mocked unit tests for both components Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
This commit is contained in:
@@ -0,0 +1,108 @@
|
||||
"""Unit tests for PostgresModelUsageStore, using a mocked psycopg2 pool.
|
||||
|
||||
Mirrors the mocking pattern in backend/tests/perception/test_postgres_event_store.py
|
||||
— no real database is needed.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from datetime import datetime, timezone
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
# Patch psycopg2 before importing the module under test.
|
||||
mock_psycopg2 = MagicMock()
|
||||
mock_psycopg2.extras = MagicMock()
|
||||
sys.modules.setdefault("psycopg2", mock_psycopg2)
|
||||
sys.modules.setdefault("psycopg2.extras", mock_psycopg2.extras)
|
||||
sys.modules.setdefault("psycopg2.pool", MagicMock())
|
||||
|
||||
from app.shared.model_usage_tracker import ModelUsageEntry
|
||||
|
||||
|
||||
def _cursor_returning(rows):
|
||||
"""Build a MagicMock standing in for a psycopg2 cursor context manager."""
|
||||
cursor = MagicMock()
|
||||
cursor.__enter__ = lambda s: s
|
||||
cursor.__exit__ = MagicMock(return_value=False)
|
||||
cursor.fetchall.return_value = rows
|
||||
return cursor
|
||||
|
||||
|
||||
@patch("app.infrastructure.storage.postgres_model_usage_store.PostgresModelUsageStore._ensure_schema")
|
||||
@patch("app.infrastructure.storage.postgres_model_usage_store.ThreadedConnectionPool")
|
||||
def test_load_all_returns_entries_keyed_by_provider_model(mock_pool_class, mock_ensure):
|
||||
"""load_all() must turn each row into a ModelUsageEntry keyed by 'provider:model'."""
|
||||
row = {
|
||||
"provider": "deepseek",
|
||||
"model": "deepseek-v4-flash",
|
||||
"total_tokens": 100,
|
||||
"prompt_tokens": 60,
|
||||
"completion_tokens": 40,
|
||||
"call_count_ok": 5,
|
||||
"call_count_error": 1,
|
||||
"last_called_at": datetime(2026, 7, 23, tzinfo=timezone.utc),
|
||||
"last_latency_ms": 250,
|
||||
"last_error": None,
|
||||
}
|
||||
mock_pool = MagicMock()
|
||||
mock_pool_class.return_value = mock_pool
|
||||
conn = MagicMock()
|
||||
conn.__enter__ = lambda s: s
|
||||
conn.__exit__ = MagicMock(return_value=False)
|
||||
conn.cursor.return_value = _cursor_returning([row])
|
||||
mock_pool.getconn.return_value = conn
|
||||
|
||||
from app.infrastructure.storage.postgres_model_usage_store import PostgresModelUsageStore
|
||||
store = PostgresModelUsageStore()
|
||||
entries = store.load_all()
|
||||
|
||||
assert "deepseek:deepseek-v4-flash" in entries
|
||||
entry = entries["deepseek:deepseek-v4-flash"]
|
||||
assert isinstance(entry, ModelUsageEntry)
|
||||
assert entry.total_tokens == 100
|
||||
assert entry.call_count_error == 1
|
||||
|
||||
|
||||
@patch("app.infrastructure.storage.postgres_model_usage_store.PostgresModelUsageStore._ensure_schema")
|
||||
@patch("app.infrastructure.storage.postgres_model_usage_store.ThreadedConnectionPool")
|
||||
def test_flush_upserts_every_entry(mock_pool_class, mock_ensure):
|
||||
"""flush() must execute one UPSERT per tracked entry and commit once."""
|
||||
mock_pool = MagicMock()
|
||||
mock_pool_class.return_value = mock_pool
|
||||
conn = MagicMock()
|
||||
conn.__enter__ = lambda s: s
|
||||
conn.__exit__ = MagicMock(return_value=False)
|
||||
cursor = MagicMock()
|
||||
cursor.__enter__ = lambda s: s
|
||||
cursor.__exit__ = MagicMock(return_value=False)
|
||||
conn.cursor.return_value = cursor
|
||||
mock_pool.getconn.return_value = conn
|
||||
|
||||
from app.infrastructure.storage.postgres_model_usage_store import PostgresModelUsageStore
|
||||
store = PostgresModelUsageStore()
|
||||
entries = {
|
||||
"deepseek:deepseek-v4-flash": ModelUsageEntry(
|
||||
provider="deepseek", model="deepseek-v4-flash", total_tokens=100, call_count_ok=5,
|
||||
),
|
||||
}
|
||||
|
||||
store.flush(entries)
|
||||
|
||||
assert cursor.execute.call_count == 1
|
||||
conn.commit.assert_called_once()
|
||||
|
||||
|
||||
@patch("app.infrastructure.storage.postgres_model_usage_store.PostgresModelUsageStore._ensure_schema")
|
||||
@patch("app.infrastructure.storage.postgres_model_usage_store.ThreadedConnectionPool")
|
||||
def test_flush_with_no_entries_does_not_touch_the_database(mock_pool_class, mock_ensure):
|
||||
"""flush({}) must be a no-op — no point opening a connection for nothing."""
|
||||
mock_pool = MagicMock()
|
||||
mock_pool_class.return_value = mock_pool
|
||||
|
||||
from app.infrastructure.storage.postgres_model_usage_store import PostgresModelUsageStore
|
||||
store = PostgresModelUsageStore()
|
||||
|
||||
store.flush({})
|
||||
|
||||
mock_pool.getconn.assert_not_called()
|
||||
@@ -77,3 +77,33 @@ def test_snapshot_returns_independent_copy():
|
||||
def test_get_model_usage_tracker_returns_singleton():
|
||||
"""get_model_usage_tracker() always returns the same process-wide instance."""
|
||||
assert get_model_usage_tracker() is get_model_usage_tracker()
|
||||
|
||||
|
||||
def test_seed_populates_registry_from_persisted_entries():
|
||||
"""seed() must bulk-load entries (e.g. from Postgres at startup) into the registry."""
|
||||
tracker = ModelUsageTracker()
|
||||
persisted = {
|
||||
"deepseek:deepseek-v4-flash": ModelUsageEntry(
|
||||
provider="deepseek", model="deepseek-v4-flash", total_tokens=500, call_count_ok=20,
|
||||
),
|
||||
}
|
||||
|
||||
tracker.seed(persisted)
|
||||
|
||||
entry = tracker.get("deepseek", "deepseek-v4-flash")
|
||||
assert entry.total_tokens == 500
|
||||
assert entry.call_count_ok == 20
|
||||
|
||||
|
||||
def test_seed_then_record_accumulates_on_top_of_seeded_value():
|
||||
"""A call recorded after seeding must add to the seeded total, not replace it."""
|
||||
tracker = ModelUsageTracker()
|
||||
tracker.seed({
|
||||
"deepseek:deepseek-v4-flash": ModelUsageEntry(
|
||||
provider="deepseek", model="deepseek-v4-flash", total_tokens=500,
|
||||
),
|
||||
})
|
||||
|
||||
tracker.record(provider="deepseek", model="deepseek-v4-flash", success=True, usage={"total_tokens": 10})
|
||||
|
||||
assert tracker.get("deepseek", "deepseek-v4-flash").total_tokens == 510
|
||||
|
||||
Reference in New Issue
Block a user