feat: record embedding call usage into ModelUsageTracker

This commit is contained in:
wangwei
2026-07-02 15:15:03 +08:00
parent 4fea159f5b
commit 41096369d3
2 changed files with 99 additions and 11 deletions
@@ -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]]: