feat: record embedding call usage into ModelUsageTracker
This commit is contained in:
@@ -3,11 +3,13 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import os
|
import os
|
||||||
|
import time
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
|
|
||||||
from app.config.settings import settings
|
from app.config.settings import settings
|
||||||
from app.domain.retrieval import EmbeddingProvider
|
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.
|
# Keep adapter behavior explicit so integration details remain easy to audit.
|
||||||
|
|
||||||
EMBEDDING_BATCH_SIZE = 8
|
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."""
|
"""Handle request for this module for the Open A I Compatible Embedding Provider instance."""
|
||||||
if not self.api_key:
|
if not self.api_key:
|
||||||
raise ValueError("缺少 EMBEDDING_API_KEY / OPENAI_API_KEY")
|
raise ValueError("缺少 EMBEDDING_API_KEY / OPENAI_API_KEY")
|
||||||
response = httpx.post(
|
start = time.time()
|
||||||
f"{self.base_url}/embeddings",
|
try:
|
||||||
headers={
|
response = httpx.post(
|
||||||
"Authorization": f"Bearer {self.api_key}",
|
f"{self.base_url}/embeddings",
|
||||||
"Content-Type": "application/json",
|
headers={
|
||||||
},
|
"Authorization": f"Bearer {self.api_key}",
|
||||||
json={"model": self.model, "input": texts},
|
"Content-Type": "application/json",
|
||||||
timeout=self.timeout,
|
},
|
||||||
)
|
json={"model": self.model, "input": texts},
|
||||||
self._raise_for_status(response, batch_size=len(texts))
|
timeout=self.timeout,
|
||||||
data = response.json()
|
)
|
||||||
|
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"])]
|
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):
|
if any(len(vector) != self.dimension for vector in vectors):
|
||||||
raise ValueError(f"embedding 维度不匹配,期望 {self.dimension}")
|
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
|
return vectors
|
||||||
|
|
||||||
def embed_texts(self, texts: list[str]) -> list[list[float]]:
|
def embed_texts(self, texts: list[str]) -> list[list[float]]:
|
||||||
|
|||||||
@@ -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:<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
|
||||||
Reference in New Issue
Block a user