Files

1187 lines
50 KiB
Markdown
Raw Permalink Normal View History

# System Status — AI Models Panel Hardening Implementation Plan
> **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking.
**Goal:** Close the three gaps left open by the already-shipped "AI Models" Status page card: streaming chat calls don't report token usage, the Cross-Encoder reranker is still disabled, and usage counters reset on every backend restart.
**Architecture:** (A1) Both provider LLM clients' `stream_chat()` generators capture a trailing OpenAI-compatible `stream_options.include_usage` chunk and return it as the generator's return value; `TrackedLLMClient.stream_chat()` retrieves that value via `StopIteration` and folds it into its existing single `tracker.record()` call. (A2) Pure config flip — the reranker code already has graceful fallback. (A3) A new `PostgresModelUsageStore` (same `CREATE TABLE IF NOT EXISTS` idiom as every other Postgres store in this codebase) is loaded once at startup to seed `ModelUsageTracker`, then flushed every 60s by a background `asyncio` task — both gated behind the existing `document_repository_backend == "postgres"` setting, so JSON-mode dev is unaffected.
**Tech Stack:** Python 3.12, FastAPI, httpx (streaming), psycopg2 (`ThreadedConnectionPool`), pytest + pytest-asyncio (`asyncio_mode = "auto"`), unittest.mock.
## Global Constraints
- Design source of truth: `docs/superpowers/specs/2026-07-23-status-model-usage-hardening-design.md`.
- All comments and docstrings in `backend/**/*.py` must be in English; every function/method needs a docstring; every file needs a module docstring + at least one meaningful `#` comment (`AGENTS.md`).
- No new business orchestration in `services/*` or `workflows/*` — this is cross-cutting observability support, same tier as `app/shared/bootstrap.py` and `app/shared/model_usage_tracker.py`.
- No migration framework — new tables use `CREATE TABLE IF NOT EXISTS` executed at first use, matching every existing Postgres store (`postgres_event_store.py`, `postgres_document_repository.py`, `postgres_document_processing_store.py`, `compliance/repository.py`, `auth/user_store.py`).
- A3 persists **current cumulative counters only** — no historical time-series table (explicit scope decision from brainstorming).
- Reuse the existing `settings.document_repository_backend` toggle (`"json"` / `"postgres"`) to gate persistence — do not add a new setting.
- Verified baseline test command (run from repo root, before any change in this plan): `python -m pytest backend/tests -q``54 passed` (1.66s). Re-run this after every task.
- Verified fast targeted command: `python -m pytest backend/tests/observability -v``18 passed` (1.01s).
---
### Task 1: DeepSeekClient — capture streaming token usage
**Files:**
- Modify: `backend/app/services/llm/deepseek_client.py`
- Test: Create `backend/tests/observability/test_stream_chat_usage_capture.py`
**Interfaces:**
- Consumes: `app.services.llm.base_client.LLMConfig`, `LLMProvider` (existing).
- Produces: `DeepSeekClient.stream_chat(...)` becomes a generator whose return value (readable via `StopIteration.value` when manually driven with `next()`) is `Optional[Dict[str, int]]` — the trailing usage dict, or `None` if the gateway never sent one. Per-chunk `yield` behavior (plain `str` content) is unchanged. Task 3 relies on this return-value contract.
- [ ] **Step 1: Write the failing test**
Create `backend/tests/observability/test_stream_chat_usage_capture.py`:
```python
"""Unit tests verifying stream_chat() captures a trailing usage-only SSE chunk.
Exercises DeepSeekClient, QwenClient, and QwenVLClient directly (not through
TrackedLLMClient) by mocking the underlying httpx.Client.stream() call — none
of these tests make a real network call.
"""
from __future__ import annotations
import json
from unittest.mock import MagicMock
from app.services.llm.base_client import LLMConfig, LLMProvider
from app.services.llm.deepseek_client import DeepSeekClient
def _sse_lines(*chunks: str, usage: dict | None = None) -> list[str]:
"""Build raw SSE 'data: ...' lines the way an OpenAI-compatible gateway sends them."""
lines = [
f'data: {json.dumps({"choices": [{"delta": {"content": c}}]})}'
for c in chunks
]
if usage is not None:
# Trailing usage-only chunk, as sent when stream_options.include_usage=true.
lines.append(f'data: {json.dumps({"choices": [], "usage": usage})}')
lines.append("data: [DONE]")
return lines
def _mock_streaming_client(lines: list[str]) -> MagicMock:
"""Build a MagicMock standing in for httpx.Client, configured for .stream()."""
fake_response = MagicMock()
fake_response.raise_for_status.return_value = None
fake_response.iter_lines.return_value = lines
stream_cm = MagicMock()
stream_cm.__enter__.return_value = fake_response
stream_cm.__exit__.return_value = False
client = MagicMock()
client.stream.return_value = stream_cm
return client
def _drain(gen):
"""Manually drive a generator, returning (yielded_chunks, stop_iteration_value)."""
chunks = []
value = None
while True:
try:
chunks.append(next(gen))
except StopIteration as stop:
value = stop.value
break
return chunks, value
def test_deepseek_stream_chat_returns_usage_from_trailing_chunk():
"""DeepSeekClient.stream_chat() must return the trailing usage dict."""
config = LLMConfig(provider=LLMProvider.DEEPSEEK, model="deepseek-v4-flash", api_key="k", base_url="http://x/v1")
client = DeepSeekClient(config)
usage = {"prompt_tokens": 5, "completion_tokens": 3, "total_tokens": 8}
client._client = _mock_streaming_client(_sse_lines("Hello", " world", usage=usage))
chunks, returned_usage = _drain(client.stream_chat([{"role": "user", "content": "hi"}]))
assert chunks == ["Hello", " world"]
assert returned_usage == usage
# The gateway must actually be asked to include usage in the stream.
sent_payload = client._client.stream.call_args.kwargs["json"]
assert sent_payload["stream_options"] == {"include_usage": True}
def test_deepseek_stream_chat_without_usage_chunk_returns_none():
"""If the gateway never sends a usage chunk, the generator returns None (unchanged behavior)."""
config = LLMConfig(provider=LLMProvider.DEEPSEEK, model="deepseek-v4-flash", api_key="k", base_url="http://x/v1")
client = DeepSeekClient(config)
client._client = _mock_streaming_client(_sse_lines("Hi"))
chunks, returned_usage = _drain(client.stream_chat([{"role": "user", "content": "hi"}]))
assert chunks == ["Hi"]
assert returned_usage is None
```
- [ ] **Step 2: Run test to verify it fails**
Run: `python -m pytest backend/tests/observability/test_stream_chat_usage_capture.py -v`
Expected: `test_deepseek_stream_chat_returns_usage_from_trailing_chunk` FAILS — `assert returned_usage == usage` fails because `returned_usage` is `None` (today's code silently drops the trailing chunk and returns nothing). `test_deepseek_stream_chat_without_usage_chunk_returns_none` PASSES already (no behavior change needed for that path).
- [ ] **Step 3: Implement — modify `backend/app/services/llm/deepseek_client.py`**
First, add `Generator` to the typing import (line 8):
```python
from typing import List, Dict, Optional, Generator
```
Then replace the existing `stream_chat` method (currently lines 127-179) with:
```python
def stream_chat(
self,
messages: List[Dict[str, str]],
max_tokens: Optional[int] = None,
temperature: Optional[float] = None,
**kwargs
) -> Generator[str, None, Optional[Dict[str, int]]]:
"""Stream chat for the Deep Seek Client instance.
Returns the trailing token-usage dict as the generator's return value
(read via StopIteration.value when manually driven with next()) when
the gateway sends one via stream_options.include_usage, else None.
"""
usage: Optional[Dict[str, int]] = None
try:
payload = {
"model": self.config.model,
"messages": messages,
"max_tokens": max_tokens or self.config.max_tokens,
"temperature": temperature or self.config.temperature,
"top_p": kwargs.get("top_p", self.config.top_p),
"stream": True,
"stream_options": {"include_usage": True}
}
with self._client.stream("POST", "/chat/completions", json=payload) as response:
response.raise_for_status()
for line in response.iter_lines():
if not line:
continue
line = line.strip()
if line.startswith(":"):
continue
if not line.startswith("data: "):
continue
data_str = line[6:]
if data_str == "[DONE]":
break
try:
import json
data = json.loads(data_str)
choices = data.get("choices", [])
if choices:
delta = choices[0].get("delta", {})
content = delta.get("content", "")
if content:
yield content
elif data.get("usage"):
# Trailing usage-only chunk — no content to yield, just capture it.
usage = data["usage"]
except json.JSONDecodeError:
continue
except httpx.HTTPStatusError as e:
logger.error(f"DeepSeek Stream API错误: {e.response.status_code}")
yield ""
except Exception as e:
logger.error(f"DeepSeek Stream调用失败: {e}")
yield ""
return usage
```
- [ ] **Step 4: Run test to verify it passes**
Run: `python -m pytest backend/tests/observability/test_stream_chat_usage_capture.py -v`
Expected: both tests PASS.
- [ ] **Step 5: Commit**
```bash
git add backend/app/services/llm/deepseek_client.py backend/tests/observability/test_stream_chat_usage_capture.py
git commit -m "feat: capture streaming token usage in DeepSeekClient.stream_chat"
```
---
### Task 2: QwenClient + QwenVLClient — capture streaming token usage
**Files:**
- Modify: `backend/app/services/llm/qwen_client.py`
- Test: Modify `backend/tests/observability/test_stream_chat_usage_capture.py` (created in Task 1)
**Interfaces:**
- Consumes: same `LLMConfig`/`LLMProvider` as Task 1.
- Produces: `QwenClient.stream_chat(...)` and `QwenVLClient.stream_chat(...)` gain the identical `Optional[Dict[str, int]]` return-value contract as `DeepSeekClient.stream_chat(...)` from Task 1. Task 3 relies on this being true for all three provider clients uniformly.
- [ ] **Step 1: Write the failing tests**
Append to `backend/tests/observability/test_stream_chat_usage_capture.py`:
```python
from app.services.llm.qwen_client import QwenClient, QwenVLClient
def test_qwen_stream_chat_returns_usage_from_trailing_chunk():
"""QwenClient.stream_chat() must return the trailing usage dict."""
config = LLMConfig(provider=LLMProvider.QWEN, model="qwen3.5-flash", api_key="k", base_url="http://x/v1")
client = QwenClient(config)
usage = {"prompt_tokens": 10, "completion_tokens": 4, "total_tokens": 14}
client._client = _mock_streaming_client(_sse_lines("Bonjour", usage=usage))
chunks, returned_usage = _drain(client.stream_chat([{"role": "user", "content": "hi"}]))
assert chunks == ["Bonjour"]
assert returned_usage == usage
sent_payload = client._client.stream.call_args.kwargs["json"]
assert sent_payload["stream_options"] == {"include_usage": True}
def test_qwen_vl_stream_chat_returns_usage_from_trailing_chunk():
"""QwenVLClient.stream_chat() must return the trailing usage dict."""
config = LLMConfig(provider=LLMProvider.QWEN_VL, model="qwen3-vl-plus", api_key="k", base_url="http://x/v1")
client = QwenVLClient(config)
usage = {"prompt_tokens": 20, "completion_tokens": 6, "total_tokens": 26}
client._client = _mock_streaming_client(_sse_lines("Describing image", usage=usage))
chunks, returned_usage = _drain(client.stream_chat([{"role": "user", "content": "describe"}]))
assert chunks == ["Describing image"]
assert returned_usage == usage
```
- [ ] **Step 2: Run tests to verify they fail**
Run: `python -m pytest backend/tests/observability/test_stream_chat_usage_capture.py -v`
Expected: the two new tests FAIL (`returned_usage` is `None` today); the two Task 1 tests still PASS.
- [ ] **Step 3: Implement — modify `backend/app/services/llm/qwen_client.py`**
`QwenClient.stream_chat` and `QwenVLClient.stream_chat` already import `Generator`, `Dict`, `Optional` (no import changes needed in this file).
Replace `QwenClient.stream_chat` (currently lines 137-184) with:
```python
def stream_chat(
self,
messages: List[Dict[str, str]],
max_tokens: Optional[int] = None,
temperature: Optional[float] = None,
**kwargs
) -> Generator[str, None, Optional[Dict[str, int]]]:
"""Stream chat for the Qwen Client instance.
Returns the trailing token-usage dict as the generator's return value
(read via StopIteration.value when manually driven with next()) when
the gateway sends one via stream_options.include_usage, else None.
"""
usage: Optional[Dict[str, int]] = None
try:
# Keep provider-specific behavior explicit so debugging stays straightforward.
payload = {
"model": self.config.model,
"messages": messages,
"max_tokens": max_tokens or self.config.max_tokens,
"temperature": temperature or self.config.temperature,
"top_p": kwargs.get("top_p", self.config.top_p),
"stream": True, # Keep provider-specific behavior explicit so debugging stays straightforward.
"stream_options": {"include_usage": True}
}
# Keep provider-specific behavior explicit so debugging stays straightforward.
with self._client.stream("POST", "/chat/completions", json=payload) as response:
for line in response.iter_lines():
if line:
line = line.strip()
# Keep provider-specific behavior explicit so debugging stays straightforward.
if line.startswith("data: "):
data_str = line[6:] # Keep provider-specific behavior explicit so debugging stays straightforward.
if data_str == "[DONE]":
break
try:
data = json.loads(data_str)
choices = data.get("choices", [])
if not choices:
if data.get("usage"):
# Trailing usage-only chunk — capture it, nothing to yield.
usage = data["usage"]
continue # Keep provider-specific behavior explicit so debugging stays straightforward.
delta = choices[0].get("delta", {})
content = delta.get("content", "")
if content:
yield content
except json.JSONDecodeError:
continue
except httpx.HTTPStatusError as e:
logger.error(f"Qwen流式API错误: {e.response.status_code}")
yield f"[ERROR: API返回错误 {e.response.status_code}]"
except Exception as e:
logger.error(f"Qwen流式调用失败: {e}")
yield f"[ERROR: {str(e)}]"
return usage
```
Replace `QwenVLClient.stream_chat` (currently lines 296-336) with:
```python
def stream_chat(
self,
messages: List[Dict[str, str]],
max_tokens: Optional[int] = None,
temperature: Optional[float] = None,
**kwargs
) -> Generator[str, None, Optional[Dict[str, int]]]:
"""Stream chat for the Qwen V L Client instance.
Returns the trailing token-usage dict as the generator's return value
(read via StopIteration.value when manually driven with next()) when
the gateway sends one via stream_options.include_usage, else None.
"""
usage: Optional[Dict[str, int]] = None
try:
payload = {
"model": self.config.model,
"messages": messages,
"max_tokens": max_tokens or self.config.max_tokens,
"temperature": temperature or self.config.temperature,
"top_p": kwargs.get("top_p", self.config.top_p),
"stream": True,
"stream_options": {"include_usage": True}
}
with self._client.stream("POST", "/chat/completions", json=payload) as response:
for line in response.iter_lines():
if line:
line = line.strip()
if line.startswith("data: "):
data_str = line[6:]
if data_str == "[DONE]":
break
try:
data = json.loads(data_str)
choices = data.get("choices", [])
if not choices:
if data.get("usage"):
# Trailing usage-only chunk — capture it, nothing to yield.
usage = data["usage"]
continue # Keep provider-specific behavior explicit so debugging stays straightforward.
delta = choices[0].get("delta", {})
content = delta.get("content", "")
if content:
yield content
except json.JSONDecodeError:
continue
except Exception as e:
logger.error(f"QwenVL流式调用失败: {e}")
yield f"[ERROR: {str(e)}]"
return usage
```
- [ ] **Step 4: Run tests to verify they pass**
Run: `python -m pytest backend/tests/observability/test_stream_chat_usage_capture.py -v`
Expected: all 4 tests PASS.
- [ ] **Step 5: Commit**
```bash
git add backend/app/services/llm/qwen_client.py backend/tests/observability/test_stream_chat_usage_capture.py
git commit -m "feat: capture streaming token usage in QwenClient and QwenVLClient"
```
---
### Task 3: TrackedLLMClient — wire streaming usage into the tracker
**Files:**
- Modify: `backend/app/services/llm/tracked_client.py`
- Test: Modify `backend/tests/observability/test_tracked_client.py`
**Interfaces:**
- Consumes: the `Optional[Dict[str, int]]` generator return-value contract from Tasks 1 and 2 (any `BaseLLMClient.stream_chat()` implementation that follows it).
- Produces: `TrackedLLMClient.stream_chat(...)` now calls `self._tracker.record(..., usage=<captured dict or None>, ...)` — the single existing call site, no new tracker methods.
- [ ] **Step 1: Write the failing test**
Append to `backend/tests/observability/test_tracked_client.py`:
```python
def test_stream_chat_records_usage_from_generator_return_value():
"""stream_chat() must forward the inner generator's returned usage dict to record()."""
inner = _make_inner()
def fake_stream(*args, **kwargs):
yield "chunk-1"
yield "chunk-2"
return {"prompt_tokens": 6, "completion_tokens": 2, "total_tokens": 8}
inner.stream_chat.side_effect = fake_stream
tracker = ModelUsageTracker()
chunks = list(TrackedLLMClient(inner, tracker).stream_chat([{"role": "user", "content": "hi"}]))
assert chunks == ["chunk-1", "chunk-2"]
entry = tracker.get("deepseek", "deepseek-v4-flash")
assert entry.total_tokens == 8
assert entry.call_count_ok == 1
```
- [ ] **Step 2: Run test to verify it fails**
Run: `python -m pytest backend/tests/observability/test_tracked_client.py -v`
Expected: `test_stream_chat_records_usage_from_generator_return_value` FAILS — `entry.total_tokens` is `0`, not `8` (today's plain `for` loop discards the generator's return value). All other existing tests in this file still PASS.
- [ ] **Step 3: Implement — modify `backend/app/services/llm/tracked_client.py`**
Replace the existing `stream_chat` method with:
```python
def stream_chat(self, messages: List[Dict[str, str]], *args: Any, **kwargs: Any):
"""Delegate to the wrapped client's stream_chat(), recording call outcome and usage.
Drives the inner generator manually (instead of a plain `for` loop) so
it can capture the generator's return value via StopIteration.value —
the trailing token-usage dict the inner client captures from a
stream_options.include_usage chunk, if the gateway sent one.
"""
start = time.time()
error: Optional[str] = None
usage: Optional[Dict[str, int]] = None
gen = self._inner.stream_chat(messages, *args, **kwargs)
try:
while True:
try:
chunk = next(gen)
except StopIteration as stop:
usage = stop.value
break
yield chunk
except Exception as exc: # noqa: BLE001 - report, then re-raise unchanged
error = str(exc)
raise
finally:
self._tracker.record(
provider=self._inner.config.provider.value,
model=self._inner.config.model,
success=error is None,
usage=usage,
latency_ms=int((time.time() - start) * 1000),
error=error,
)
```
- [ ] **Step 4: Run tests to verify they pass**
Run: `python -m pytest backend/tests/observability/test_tracked_client.py -v`
Expected: all 6 tests PASS (5 existing + 1 new).
- [ ] **Step 5: Run the full targeted suite**
Run: `python -m pytest backend/tests/observability -v`
Expected: all tests PASS (18 existing + 5 new from Tasks 1-3 = 23 total... exact count will depend on final additions, all green).
- [ ] **Step 6: Commit**
```bash
git add backend/app/services/llm/tracked_client.py backend/tests/observability/test_tracked_client.py
git commit -m "feat: record streaming token usage in TrackedLLMClient.stream_chat"
```
---
### Task 4: Enable the Cross-Encoder reranker
**Files:**
- Modify: `.env` (repo root)
**Interfaces:**
- Consumes: nothing new — `OpenAICompatibleReranker` (`backend/app/infrastructure/vectorstore/cross_encoder_reranker.py`, unchanged) already reads `settings.reranker_enabled`.
- Produces: nothing new — this is a pure configuration change, exercised by existing code.
- [ ] **Step 1: Flip the flag**
In `.env` at the repo root, change:
```
RERANKER_ENABLED=false
```
to:
```
RERANKER_ENABLED=true
```
- [ ] **Step 2: Verify via the already-shipped ping endpoint**
Start the backend locally (`dev.bat start api --foreground` or `PYTHONPATH=backend uvicorn app.main:app --host 0.0.0.0 --port 8000 --reload`), then:
```bash
curl -X POST http://localhost:8000/api/v1/status/models/ping
```
Expected: the JSON response's `reranker` entry has `"enabled": true` and `"status"` is `"ok"` (gateway responded) — not `"error"`. If it is `"error"`, the gateway does not currently support the configured `/rerank` endpoint; revert `.env` to `RERANKER_ENABLED=false` and stop here (do not proceed to commit) — this needs a separate follow-up to fix gateway connectivity, out of scope for this plan.
- [ ] **Step 3: Commit**
```bash
git add .env
git commit -m "chore: enable Cross-Encoder reranker (verified via /status/models/ping)"
```
---
### Task 5: PostgresModelUsageStore — persistence primitive
**Files:**
- Modify: `backend/app/shared/model_usage_tracker.py` (add `seed()`)
- Create: `backend/app/infrastructure/storage/postgres_model_usage_store.py`
- Test: Modify `backend/tests/observability/test_model_usage_tracker.py` (for `seed()`)
- Test: Create `backend/tests/observability/test_model_usage_persistence.py`
**Interfaces:**
- Consumes: `app.shared.model_usage_tracker.ModelUsageEntry` (existing dataclass: `provider: str, model: str, total_tokens: int = 0, prompt_tokens: int = 0, completion_tokens: int = 0, call_count_ok: int = 0, call_count_error: int = 0, last_called_at: datetime | None = None, last_latency_ms: int | None = None, last_error: str | None = None`).
- Produces:
- `ModelUsageTracker.seed(entries: dict[str, ModelUsageEntry]) -> None` — bulk-loads persisted entries at startup.
- `PostgresModelUsageStore.load_all() -> dict[str, ModelUsageEntry]` — reads the whole table, keyed `"{provider}:{model}"`.
- `PostgresModelUsageStore.flush(entries: dict[str, ModelUsageEntry]) -> None` — upserts every entry's current snapshot.
Task 6 consumes both of the above.
- [ ] **Step 1: Write the failing test for `ModelUsageTracker.seed()`**
Append to `backend/tests/observability/test_model_usage_tracker.py`:
```python
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
```
(Check the top of the test file already imports `ModelUsageEntry` — if not, add `from app.shared.model_usage_tracker import ModelUsageEntry, ModelUsageTracker` to the imports.)
- [ ] **Step 2: Run test to verify it fails**
Run: `python -m pytest backend/tests/observability/test_model_usage_tracker.py -v`
Expected: both new tests FAIL with `AttributeError: 'ModelUsageTracker' object has no attribute 'seed'`.
- [ ] **Step 3: Implement `seed()` — modify `backend/app/shared/model_usage_tracker.py`**
Add this method to the `ModelUsageTracker` class, directly after `record()` and before `snapshot()`:
```python
def seed(self, entries: dict[str, ModelUsageEntry]) -> None:
"""Bulk-load persisted entries (called once at startup, before any traffic).
Unlike record(), this replaces entries wholesale rather than
accumulating deltas — it exists to restore counters saved by a
previous process run, not to record a new call.
"""
with self._lock:
self._entries.update(entries)
```
- [ ] **Step 4: Run test to verify it passes**
Run: `python -m pytest backend/tests/observability/test_model_usage_tracker.py -v`
Expected: all tests PASS.
- [ ] **Step 5: Write the failing tests for `PostgresModelUsageStore`**
Create `backend/tests/observability/test_model_usage_persistence.py`:
```python
"""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
def _make_store_with_pool(mock_pool):
"""Construct PostgresModelUsageStore with its connection pool replaced by a mock."""
with patch("psycopg2.pool.ThreadedConnectionPool", return_value=mock_pool):
with patch(
"app.infrastructure.storage.postgres_model_usage_store.PostgresModelUsageStore._ensure_schema"
):
from app.infrastructure.storage.postgres_model_usage_store import PostgresModelUsageStore
return PostgresModelUsageStore()
def test_load_all_returns_entries_keyed_by_provider_model():
"""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()
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
store = _make_store_with_pool(mock_pool)
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
def test_flush_upserts_every_entry():
"""flush() must execute one UPSERT per tracked entry and commit once."""
mock_pool = MagicMock()
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
store = _make_store_with_pool(mock_pool)
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()
def test_flush_with_no_entries_does_not_touch_the_database():
"""flush({}) must be a no-op — no point opening a connection for nothing."""
mock_pool = MagicMock()
store = _make_store_with_pool(mock_pool)
store.flush({})
mock_pool.getconn.assert_not_called()
```
- [ ] **Step 6: Run tests to verify they fail**
Run: `python -m pytest backend/tests/observability/test_model_usage_persistence.py -v`
Expected: FAIL with `ModuleNotFoundError: No module named 'app.infrastructure.storage.postgres_model_usage_store'`.
- [ ] **Step 7: Implement — create `backend/app/infrastructure/storage/postgres_model_usage_store.py`**
```python
"""Postgres-backed persistence for cumulative AI model usage counters.
Keeps ModelUsageTracker (an in-memory, process-lifetime-only registry defined
in app/shared/model_usage_tracker.py) from losing its counters on every
backend restart. This store only ever persists the *current cumulative
snapshot* per provider+model — not a historical time-series log — matching
the "durable counters" scope decided in
docs/superpowers/specs/2026-07-23-status-model-usage-hardening-design.md.
"""
from __future__ import annotations
from contextlib import contextmanager
import psycopg2
import psycopg2.extras
from psycopg2.pool import ThreadedConnectionPool
from app.config.settings import settings
from app.shared.model_usage_tracker import ModelUsageEntry
# Table creation follows the same CREATE TABLE IF NOT EXISTS idiom used by
# every other Postgres store in this codebase — no migration framework.
_CREATE_TABLE = """
CREATE TABLE IF NOT EXISTS model_usage_stats (
provider VARCHAR(64) NOT NULL,
model VARCHAR(128) NOT NULL,
total_tokens BIGINT NOT NULL DEFAULT 0,
prompt_tokens BIGINT NOT NULL DEFAULT 0,
completion_tokens BIGINT NOT NULL DEFAULT 0,
call_count_ok BIGINT NOT NULL DEFAULT 0,
call_count_error BIGINT NOT NULL DEFAULT 0,
last_called_at TIMESTAMPTZ,
last_latency_ms INTEGER,
last_error TEXT,
updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
PRIMARY KEY (provider, model)
);
"""
_UPSERT = """
INSERT INTO model_usage_stats
(provider, model, total_tokens, prompt_tokens, completion_tokens,
call_count_ok, call_count_error, last_called_at, last_latency_ms, last_error, updated_at)
VALUES
(%(provider)s, %(model)s, %(total_tokens)s, %(prompt_tokens)s, %(completion_tokens)s,
%(call_count_ok)s, %(call_count_error)s, %(last_called_at)s, %(last_latency_ms)s, %(last_error)s, NOW())
ON CONFLICT (provider, model) DO UPDATE SET
total_tokens = EXCLUDED.total_tokens,
prompt_tokens = EXCLUDED.prompt_tokens,
completion_tokens = EXCLUDED.completion_tokens,
call_count_ok = EXCLUDED.call_count_ok,
call_count_error = EXCLUDED.call_count_error,
last_called_at = EXCLUDED.last_called_at,
last_latency_ms = EXCLUDED.last_latency_ms,
last_error = EXCLUDED.last_error,
updated_at = NOW();
"""
class PostgresModelUsageStore:
"""Load and flush ModelUsageTracker snapshots to/from a Postgres table."""
def __init__(self) -> None:
"""Open a small connection pool and ensure the table exists."""
self._pool = ThreadedConnectionPool(
minconn=1,
maxconn=3,
host=settings.postgres_host,
port=settings.postgres_port,
user=settings.postgres_user,
password=settings.postgres_password,
dbname=settings.postgres_db,
)
self._ensure_schema()
def _ensure_schema(self) -> None:
"""Create the model_usage_stats table if it does not already exist."""
with self._conn() as conn:
with conn.cursor() as cur:
cur.execute(_CREATE_TABLE)
conn.commit()
@contextmanager
def _conn(self):
"""Borrow a pooled connection and always return it, even on error."""
conn = self._pool.getconn()
try:
yield conn
finally:
self._pool.putconn(conn)
def load_all(self) -> dict[str, ModelUsageEntry]:
"""Return every persisted row as {"provider:model": ModelUsageEntry}."""
with self._conn() as conn:
with conn.cursor(cursor_factory=psycopg2.extras.RealDictCursor) as cur:
cur.execute("SELECT * FROM model_usage_stats")
rows = cur.fetchall()
entries: dict[str, ModelUsageEntry] = {}
for row in rows:
entry = ModelUsageEntry(
provider=row["provider"],
model=row["model"],
total_tokens=row["total_tokens"],
prompt_tokens=row["prompt_tokens"],
completion_tokens=row["completion_tokens"],
call_count_ok=row["call_count_ok"],
call_count_error=row["call_count_error"],
last_called_at=row["last_called_at"],
last_latency_ms=row["last_latency_ms"],
last_error=row["last_error"],
)
entries[f"{entry.provider}:{entry.model}"] = entry
return entries
def flush(self, entries: dict[str, ModelUsageEntry]) -> None:
"""Upsert the current cumulative snapshot of every tracked entry.
A no-op for an empty snapshot — avoids opening a connection for nothing
(e.g. before any LLM/embedding/reranker call has happened yet).
"""
if not entries:
return
with self._conn() as conn:
with conn.cursor() as cur:
for entry in entries.values():
cur.execute(
_UPSERT,
{
"provider": entry.provider,
"model": entry.model,
"total_tokens": entry.total_tokens,
"prompt_tokens": entry.prompt_tokens,
"completion_tokens": entry.completion_tokens,
"call_count_ok": entry.call_count_ok,
"call_count_error": entry.call_count_error,
"last_called_at": entry.last_called_at,
"last_latency_ms": entry.last_latency_ms,
"last_error": entry.last_error,
},
)
conn.commit()
```
- [ ] **Step 8: Run tests to verify they pass**
Run: `python -m pytest backend/tests/observability/test_model_usage_persistence.py -v`
Expected: all 3 tests PASS.
- [ ] **Step 9: Run the full targeted suite**
Run: `python -m pytest backend/tests/observability -v`
Expected: all tests PASS.
- [ ] **Step 10: Commit**
```bash
git add backend/app/shared/model_usage_tracker.py backend/app/infrastructure/storage/postgres_model_usage_store.py backend/tests/observability/test_model_usage_tracker.py backend/tests/observability/test_model_usage_persistence.py
git commit -m "feat: add PostgresModelUsageStore and ModelUsageTracker.seed()"
```
---
### Task 6: Wire persistence into app startup/shutdown
**Files:**
- Modify: `backend/app/shared/bootstrap.py`
- Test: Create `backend/tests/observability/test_model_usage_bootstrap.py`
**Interfaces:**
- Consumes: `PostgresModelUsageStore.load_all()` / `.flush()` and `ModelUsageTracker.seed()` / `.snapshot()` (all from Task 5); `settings.document_repository_backend` (existing).
- Produces: `get_model_usage_store()` (new, `@lru_cache`, mirrors `get_parse_artifact_store()`), `_start_model_usage_persistence()` and `_stop_model_usage_persistence()` (new, called from the existing `preload_runtime_dependencies()` / `cleanup_runtime_dependencies()`, which are themselves already wired into `backend/app/api/main.py`'s `lifespan()` — no change needed to `main.py`).
- [ ] **Step 1: Write the failing tests**
Create `backend/tests/observability/test_model_usage_bootstrap.py`:
```python
"""Unit tests for the model-usage persistence wiring in app.shared.bootstrap.
get_model_usage_store()'s settings-gating is tested the same way
tests/test_reranker_bootstrap.py tests get_reranker() — by patching
"app.shared.bootstrap.settings" wholesale, matching this codebase's
established convention for testing @lru_cache settings-gated factories.
The remaining tests isolate _start_model_usage_persistence() /
_stop_model_usage_persistence() from get_model_usage_store() entirely (via
monkeypatch on the module-level function), so no real database or event loop
is needed anywhere in this file — asyncio.create_task itself is also mocked.
"""
from __future__ import annotations
import sys
from unittest.mock import MagicMock, patch
# Patch psycopg2 before importing anything that transitively imports it, in
# case this file is collected before test_model_usage_persistence.py.
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 import bootstrap
from app.shared.model_usage_tracker import ModelUsageEntry, ModelUsageTracker
def test_get_model_usage_store_returns_none_when_not_postgres_backend():
"""get_model_usage_store() must be None unless document_repository_backend == 'postgres'."""
bootstrap.get_model_usage_store.cache_clear()
with patch("app.shared.bootstrap.settings") as mock_settings:
mock_settings.document_repository_backend = "json"
result = bootstrap.get_model_usage_store()
bootstrap.get_model_usage_store.cache_clear()
assert result is None
def test_get_model_usage_store_returns_instance_when_postgres_backend():
"""get_model_usage_store() must return a PostgresModelUsageStore when enabled.
ThreadedConnectionPool is mocked so no real connection is attempted; the
postgres_host/port/user/password/db values PostgresModelUsageStore reads
come from app.config.settings.settings directly (not from the
app.shared.bootstrap.settings reference mocked below), so they don't need
to be set here — only document_repository_backend gates this factory.
"""
bootstrap.get_model_usage_store.cache_clear()
with patch("psycopg2.pool.ThreadedConnectionPool"), \
patch(
"app.infrastructure.storage.postgres_model_usage_store.PostgresModelUsageStore._ensure_schema"
), \
patch("app.shared.bootstrap.settings") as mock_settings:
mock_settings.document_repository_backend = "postgres"
result = bootstrap.get_model_usage_store()
bootstrap.get_model_usage_store.cache_clear()
from app.infrastructure.storage.postgres_model_usage_store import PostgresModelUsageStore
assert isinstance(result, PostgresModelUsageStore)
def test_start_model_usage_persistence_seeds_tracker_and_starts_flush_loop(monkeypatch):
"""When a store is available, startup must seed the tracker and schedule the flush task."""
fake_store = MagicMock()
fake_store.load_all.return_value = {
"deepseek:deepseek-v4-flash": ModelUsageEntry(
provider="deepseek", model="deepseek-v4-flash", total_tokens=99,
),
}
tracker = ModelUsageTracker()
monkeypatch.setattr(bootstrap, "get_model_usage_store", lambda: fake_store)
monkeypatch.setattr(bootstrap, "get_model_usage_tracker", lambda: tracker)
with patch("asyncio.create_task") as mock_create_task:
bootstrap._start_model_usage_persistence()
# Close the coroutine object passed to the mock so pytest doesn't warn
# about "coroutine was never awaited" — it was never meant to run here.
mock_create_task.call_args[0][0].close()
assert tracker.get("deepseek", "deepseek-v4-flash").total_tokens == 99
mock_create_task.assert_called_once()
bootstrap._stop_model_usage_persistence() # reset the module-level task handle
def test_start_model_usage_persistence_is_a_no_op_without_a_store(monkeypatch):
"""No store configured (json backend) — startup must not touch asyncio or the tracker."""
monkeypatch.setattr(bootstrap, "get_model_usage_store", lambda: None)
with patch("asyncio.create_task") as mock_create_task:
bootstrap._start_model_usage_persistence()
mock_create_task.assert_not_called()
def test_stop_model_usage_persistence_cancels_task_and_flushes(monkeypatch):
"""Shutdown must cancel the running flush task and perform one final flush."""
fake_store = MagicMock()
monkeypatch.setattr(bootstrap, "get_model_usage_store", lambda: fake_store)
fake_task = MagicMock()
bootstrap._model_usage_flush_task = fake_task
bootstrap._stop_model_usage_persistence()
fake_task.cancel.assert_called_once()
fake_store.flush.assert_called_once()
assert bootstrap._model_usage_flush_task is None
```
- [ ] **Step 2: Run tests to verify they fail**
Run: `python -m pytest backend/tests/observability/test_model_usage_bootstrap.py -v`
Expected: FAIL with `AttributeError: module 'app.shared.bootstrap' has no attribute 'get_model_usage_store'` (and similar for the other new names).
- [ ] **Step 3: Implement — modify `backend/app/shared/bootstrap.py`**
Add these imports near the top (alongside the existing `from functools import lru_cache` and other stdlib imports):
```python
import asyncio
from loguru import logger
```
Add this import alongside the other `from app.infrastructure.storage...` imports:
```python
from app.infrastructure.storage.postgres_model_usage_store import PostgresModelUsageStore
```
Add this import alongside other `app.shared` imports (or near the top with the other `app.*` imports):
```python
from app.shared.model_usage_tracker import get_model_usage_tracker
```
Add this factory function directly after the existing `get_parse_artifact_store()` (which returns `PostgresParseArtifactStore()` or `None`):
```python
@lru_cache
def get_model_usage_store():
"""Return the Postgres model-usage store, or None when postgres backend is not enabled."""
if settings.document_repository_backend == "postgres":
return PostgresModelUsageStore()
return None
```
Replace the existing `preload_runtime_dependencies()` / `cleanup_runtime_dependencies()` pair with:
```python
def preload_runtime_dependencies() -> None:
"""Warm dependencies that are safe and useful to preload during startup."""
LLMFactory.preload_clients(["qwen", "deepseek"])
_start_model_usage_persistence()
def cleanup_runtime_dependencies() -> None:
"""Release runtime dependencies that expose explicit cleanup hooks."""
LLMFactory.cleanup()
_stop_model_usage_persistence()
_model_usage_flush_task: "asyncio.Task | None" = None
def _start_model_usage_persistence() -> None:
"""Seed ModelUsageTracker from Postgres and start its periodic flush loop.
No-op when document_repository_backend != "postgres" — ModelUsageTracker
then keeps behaving exactly as it always has: purely in-memory, reset on
every restart. Never raises: persistence must not block app startup.
"""
global _model_usage_flush_task
try:
store = get_model_usage_store()
except Exception as exc: # noqa: BLE001 - persistence must never block startup
logger.warning("Failed to initialize model usage persistence: {}", exc)
return
if store is None:
return
tracker = get_model_usage_tracker()
try:
tracker.seed(store.load_all())
except Exception as exc: # noqa: BLE001 - a bad load must not block startup
logger.warning("Failed to load persisted model usage stats: {}", exc)
async def _flush_loop() -> None:
"""Snapshot the tracker into Postgres every 60 seconds until cancelled."""
while True:
await asyncio.sleep(60)
try:
store.flush(tracker.snapshot())
except Exception as exc: # noqa: BLE001 - one bad cycle must not kill the loop
logger.warning("Failed to flush model usage stats: {}", exc)
_model_usage_flush_task = asyncio.create_task(_flush_loop())
def _stop_model_usage_persistence() -> None:
"""Cancel the periodic flush task and perform one best-effort final flush."""
global _model_usage_flush_task
if _model_usage_flush_task is not None:
_model_usage_flush_task.cancel()
_model_usage_flush_task = None
try:
store = get_model_usage_store()
except Exception as exc: # noqa: BLE001 - shutdown must not crash on this
logger.warning("Failed to access model usage store during shutdown: {}", exc)
return
if store is None:
return
try:
store.flush(get_model_usage_tracker().snapshot())
except Exception as exc: # noqa: BLE001 - shutdown must not crash on a flush failure
logger.warning("Failed final model usage flush: {}", exc)
```
- [ ] **Step 4: Run tests to verify they pass**
Run: `python -m pytest backend/tests/observability/test_model_usage_bootstrap.py -v`
Expected: all 6 tests PASS.
- [ ] **Step 5: Run the full backend test suite**
Run: `python -m pytest backend/tests -q`
Expected: all tests PASS (baseline was `54 passed`; this task's net new tests plus Tasks 1/2/3/5's net new tests should all be green — no regressions in the pre-existing 54).
- [ ] **Step 6: Manual smoke check (requires a reachable Postgres with `DOCUMENT_REPOSITORY_BACKEND=postgres` in `.env`)**
Start the backend, confirm in the logs there is no `Failed to initialize model usage persistence` warning, then:
```bash
curl http://localhost:8000/api/v1/status/models
```
Expected: same response shape as before this plan (this feature is transparent to the API contract) — the numbers should now survive a backend restart instead of resetting to zero. If Postgres is not reachable in your environment, skip this step — the automated tests already cover the logic in isolation.
- [ ] **Step 7: Commit**
```bash
git add backend/app/shared/bootstrap.py backend/tests/observability/test_model_usage_bootstrap.py
git commit -m "feat: seed and periodically persist model usage stats to Postgres"
```
---
## Self-Review Notes
- **Spec coverage:** A1 → Tasks 1-3. A2 → Task 4. A3 → Tasks 5-6. All three goals from the design doc have corresponding tasks; all "Out of Scope" items (cost estimation, per-session breakdown, time-series charts, Langfuse/Ragas) are correctly left untouched.
- **Placeholder scan:** no TBD/TODO; every step has complete, runnable code and exact commands.
- **Type consistency:** `Optional[Dict[str, int]]` return-value contract is named identically across Tasks 1, 2, and 3's Interfaces sections and code. `ModelUsageEntry` field names in `postgres_model_usage_store.py` (Task 5) match the dataclass exactly as defined in `backend/app/shared/model_usage_tracker.py`. `get_model_usage_store()` / `get_model_usage_tracker()` names are used identically in Tasks 5 and 6.