66 lines
2.4 KiB
Python
66 lines
2.4 KiB
Python
"""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:<model>'."""
|
||
|
|
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
|