From 41096369d3b1b1ba251525c9247d26a0d599638e Mon Sep 17 00:00:00 2001 From: wangwei Date: Thu, 2 Jul 2026 15:15:03 +0800 Subject: [PATCH] feat: record embedding call usage into ModelUsageTracker --- .../openai_compatible_embedding_provider.py | 45 +++++++++---- .../test_embedding_usage_tracking.py | 65 +++++++++++++++++++ 2 files changed, 99 insertions(+), 11 deletions(-) create mode 100644 backend/tests/observability/test_embedding_usage_tracking.py diff --git a/backend/app/infrastructure/embedding/openai_compatible_embedding_provider.py b/backend/app/infrastructure/embedding/openai_compatible_embedding_provider.py index 2b308dc..09c7537 100644 --- a/backend/app/infrastructure/embedding/openai_compatible_embedding_provider.py +++ b/backend/app/infrastructure/embedding/openai_compatible_embedding_provider.py @@ -3,11 +3,13 @@ from __future__ import annotations import os +import time import httpx from app.config.settings import settings from app.domain.retrieval import EmbeddingProvider +from app.shared.model_usage_tracker import get_model_usage_tracker # Keep adapter behavior explicit so integration details remain easy to audit. EMBEDDING_BATCH_SIZE = 8 @@ -45,20 +47,41 @@ class OpenAICompatibleEmbeddingProvider(EmbeddingProvider): """Handle request for this module for the Open A I Compatible Embedding Provider instance.""" if not self.api_key: raise ValueError("缺少 EMBEDDING_API_KEY / OPENAI_API_KEY") - response = httpx.post( - f"{self.base_url}/embeddings", - headers={ - "Authorization": f"Bearer {self.api_key}", - "Content-Type": "application/json", - }, - json={"model": self.model, "input": texts}, - timeout=self.timeout, - ) - self._raise_for_status(response, batch_size=len(texts)) - data = response.json() + start = time.time() + try: + response = httpx.post( + f"{self.base_url}/embeddings", + headers={ + "Authorization": f"Bearer {self.api_key}", + "Content-Type": "application/json", + }, + json={"model": self.model, "input": texts}, + timeout=self.timeout, + ) + self._raise_for_status(response, batch_size=len(texts)) + data = response.json() + except Exception as exc: + # Record the failed call so the Status page can show it as an error, + # then re-raise unchanged so existing callers keep their current behavior. + get_model_usage_tracker().record( + provider="embedding", + model=self.model, + success=False, + latency_ms=int((time.time() - start) * 1000), + error=str(exc), + ) + raise vectors = [item["embedding"] for item in sorted(data.get("data", []), key=lambda item: item["index"])] if any(len(vector) != self.dimension for vector in vectors): raise ValueError(f"embedding 维度不匹配,期望 {self.dimension}") + # Record token usage from the OpenAI-compatible response, e.g. {"total_tokens": N}. + get_model_usage_tracker().record( + provider="embedding", + model=self.model, + success=True, + usage=data.get("usage", {}), + latency_ms=int((time.time() - start) * 1000), + ) return vectors def embed_texts(self, texts: list[str]) -> list[list[float]]: diff --git a/backend/tests/observability/test_embedding_usage_tracking.py b/backend/tests/observability/test_embedding_usage_tracking.py new file mode 100644 index 0000000..f66f0e0 --- /dev/null +++ b/backend/tests/observability/test_embedding_usage_tracking.py @@ -0,0 +1,65 @@ +"""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