51 lines
2.1 KiB
Python
51 lines
2.1 KiB
Python
"""Verifies OpenAICompatibleReranker records call outcome (no tokens) into ModelUsageTracker."""
|
|
|
|
from __future__ import annotations
|
|
|
|
from unittest.mock import patch
|
|
|
|
import pytest
|
|
|
|
from app.domain.retrieval import RetrievedChunk
|
|
from app.infrastructure.vectorstore.cross_encoder_reranker import OpenAICompatibleReranker
|
|
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 _chunk(chunk_id: str, text: str) -> RetrievedChunk:
|
|
"""Build a minimal RetrievedChunk for reranker tests."""
|
|
return RetrievedChunk(chunk_id=chunk_id, doc_id="doc-1", doc_title="Doc", text=text, score=0.0)
|
|
|
|
|
|
def test_successful_rerank_records_call_without_tokens():
|
|
"""A successful rerank() call is recorded with call_count_ok but zero tokens."""
|
|
reranker = OpenAICompatibleReranker(base_url="http://example.test", model="bge-reranker-v2-m3")
|
|
with patch.object(reranker, "_call_reranker", return_value=[0.9, 0.1]):
|
|
result = reranker.rerank("query", [_chunk("c1", "a"), _chunk("c2", "b")], top_k=2)
|
|
|
|
assert len(result) == 2
|
|
entry = get_model_usage_tracker().get("reranker", "bge-reranker-v2-m3")
|
|
assert entry is not None
|
|
assert entry.call_count_ok == 1
|
|
assert entry.total_tokens == 0
|
|
|
|
|
|
def test_failed_rerank_records_error_and_falls_back():
|
|
"""A rerank() call that raises internally is recorded as an error but still returns a fallback list."""
|
|
reranker = OpenAICompatibleReranker(base_url="http://example.test", model="bge-reranker-v2-m3")
|
|
with patch.object(reranker, "_call_reranker", side_effect=RuntimeError("gateway down")):
|
|
result = reranker.rerank("query", [_chunk("c1", "a")], top_k=1)
|
|
|
|
assert len(result) == 1 # existing fallback behavior: original order, unscored
|
|
entry = get_model_usage_tracker().get("reranker", "bge-reranker-v2-m3")
|
|
assert entry is not None
|
|
assert entry.call_count_error == 1
|
|
assert entry.status == "error"
|