main-ruqi #1

Merged
wangwei merged 17 commits from main-ruqi into main 2026-07-02 22:05:17 +08:00
2 changed files with 99 additions and 11 deletions
Showing only changes of commit 41096369d3 - Show all commits
@@ -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