"""Verifies OpenAICompatibleEmbeddingProvider records usage into ModelUsageTracker.""" from __future__ import annotations from unittest.mock import MagicMock, patch import httpx import pytest from app.infrastructure.embedding.openai_compatible_embedding_provider import OpenAICompatibleEmbeddingProvider from app.shared.model_usage_tracker import get_model_usage_tracker @pytest.fixture(autouse=True) def _reset_tracker(): """Clear the process-wide tracker before and after each test in this file.""" get_model_usage_tracker()._entries.clear() yield get_model_usage_tracker()._entries.clear() def _fake_response(usage: dict) -> MagicMock: """Build a fake httpx.Response-like object for a successful embeddings call.""" resp = MagicMock(spec=httpx.Response) resp.raise_for_status.return_value = None resp.json.return_value = { "data": [{"index": 0, "embedding": [0.1] * 1024}], "usage": usage, } return resp def test_successful_embed_records_usage(): """A successful embeddings call must record token usage under 'embedding:'.""" provider = OpenAICompatibleEmbeddingProvider() provider.api_key = "test-key" with patch("httpx.post", return_value=_fake_response({"prompt_tokens": 3, "total_tokens": 3})): provider.embed_query("hello") entry = get_model_usage_tracker().get("embedding", provider.model) assert entry is not None assert entry.total_tokens == 3 assert entry.status == "ok" def test_failed_embed_records_error(): """An HTTP error from the embeddings endpoint must be recorded as a failure, then re-raised.""" provider = OpenAICompatibleEmbeddingProvider() provider.api_key = "test-key" failing_response = MagicMock(spec=httpx.Response) failing_response.status_code = 500 failing_response.text = "boom" failing_response.request = MagicMock() failing_response.request.url = "http://example.com/embeddings" failing_response.raise_for_status.side_effect = httpx.HTTPStatusError( "boom", request=failing_response.request, response=failing_response ) with patch("httpx.post", return_value=failing_response): with pytest.raises(httpx.HTTPStatusError): provider.embed_query("hello") entry = get_model_usage_tracker().get("embedding", provider.model) assert entry is not None assert entry.status == "error" assert entry.call_count_error == 1