135 lines
5.1 KiB
Python
135 lines
5.1 KiB
Python
"""Tests that build_models resolves separate clients for judge and embedding models."""
|
|||
|
|
|
||
|
|
from unittest.mock import MagicMock, patch, call
|
||
|
|
from rag_eval.metrics.factory import build_models
|
||
|
|
from rag_eval.settings import EvaluationSettings
|
||
|
|
|
||
|
|
|
||
|
|
def _settings():
|
||
|
|
return EvaluationSettings(_env_file=None)
|
||
|
|
|
||
|
|
|
||
|
|
def _patch_profiles(profiles_by_model: dict):
|
||
|
|
"""Patch profile_manager.list_all() to return fake profiles."""
|
||
|
|
fake_profiles = []
|
||
|
|
for model, base_url in profiles_by_model.items():
|
||
|
|
p = MagicMock()
|
||
|
|
p.model = model
|
||
|
|
p.base_url = base_url
|
||
|
|
p.api_key = "sk-test"
|
||
|
|
p.timeout_seconds = 30
|
||
|
|
fake_profiles.append(p)
|
||
|
|
return fake_profiles
|
||
|
|
|
||
|
|
|
||
|
|
def test_build_models_uses_separate_clients_for_different_profiles(monkeypatch):
|
||
|
|
"""When judge and embedding models have different profiles, two AsyncOpenAI clients are created."""
|
||
|
|
created_clients = []
|
||
|
|
|
||
|
|
class _FakeClient:
|
||
|
|
def __init__(self, **kwargs):
|
||
|
|
created_clients.append(kwargs.copy())
|
||
|
|
|
||
|
|
fake_profiles = _patch_profiles({
|
||
|
|
"gpt-5": "http://llm-gateway/v1",
|
||
|
|
"Qwen3-Embedding-4B": "https://api.siliconflow.cn/v1",
|
||
|
|
})
|
||
|
|
|
||
|
|
import webapp.services.profile_manager as pm_mod
|
||
|
|
monkeypatch.setattr(pm_mod.profile_manager, "list_all", lambda: fake_profiles)
|
||
|
|
|
||
|
|
from ragas.llms.base import InstructorBaseRagasLLM
|
||
|
|
from ragas.embeddings.base import BaseRagasEmbedding
|
||
|
|
|
||
|
|
with patch("rag_eval.metrics.factory.AsyncOpenAI", side_effect=_FakeClient), \
|
||
|
|
patch("rag_eval.metrics.factory.llm_factory", return_value=MagicMock(spec=InstructorBaseRagasLLM)), \
|
||
|
|
patch("rag_eval.metrics.factory.embedding_factory", return_value=MagicMock(spec=BaseRagasEmbedding)):
|
||
|
|
build_models("gpt-5", "Qwen3-Embedding-4B", _settings())
|
||
|
|
|
||
|
|
# Two distinct clients should have been created with different base_urls.
|
||
|
|
assert len(created_clients) == 2
|
||
|
|
base_urls = {c["base_url"] for c in created_clients}
|
||
|
|
assert "http://llm-gateway/v1" in base_urls
|
||
|
|
assert "https://api.siliconflow.cn/v1" in base_urls
|
||
|
|
|
||
|
|
|
||
|
|
def test_build_models_shares_client_when_same_profile(monkeypatch):
|
||
|
|
"""When judge and embedding models resolve to the same settings, only one client is created."""
|
||
|
|
created_clients = []
|
||
|
|
|
||
|
|
class _FakeClient:
|
||
|
|
def __init__(self, **kwargs):
|
||
|
|
created_clients.append(kwargs.copy())
|
||
|
|
|
||
|
|
# Both models map to the same gateway.
|
||
|
|
fake_profiles = _patch_profiles({
|
||
|
|
"gpt-5": "http://same-gateway/v1",
|
||
|
|
"text-embedding-3-small": "http://same-gateway/v1",
|
||
|
|
})
|
||
|
|
|
||
|
|
import webapp.services.profile_manager as pm_mod
|
||
|
|
monkeypatch.setattr(pm_mod.profile_manager, "list_all", lambda: fake_profiles)
|
||
|
|
|
||
|
|
from ragas.llms.base import InstructorBaseRagasLLM
|
||
|
|
from ragas.embeddings.base import BaseRagasEmbedding
|
||
|
|
|
||
|
|
with patch("rag_eval.metrics.factory.AsyncOpenAI", side_effect=_FakeClient), \
|
||
|
|
patch("rag_eval.metrics.factory.llm_factory", return_value=MagicMock(spec=InstructorBaseRagasLLM)), \
|
||
|
|
patch("rag_eval.metrics.factory.embedding_factory", return_value=MagicMock(spec=BaseRagasEmbedding)):
|
||
|
|
build_models("gpt-5", "text-embedding-3-small", _settings())
|
||
|
|
|
||
|
|
# Same settings → only one client needed.
|
||
|
|
assert len(created_clients) == 1
|
||
|
|
|
||
|
|
|
||
|
|
def test_embedding_factory_receives_embedding_client(monkeypatch):
|
||
|
|
"""embedding_factory is called with the client resolved from the embedding model's profile."""
|
||
|
|
emb_client_ref = []
|
||
|
|
llm_client_ref = []
|
||
|
|
|
||
|
|
class _LLMClient:
|
||
|
|
label = "llm"
|
||
|
|
def __init__(self, **kwargs):
|
||
|
|
llm_client_ref.append(self)
|
||
|
|
|
||
|
|
class _EmbClient:
|
||
|
|
label = "emb"
|
||
|
|
def __init__(self, **kwargs):
|
||
|
|
emb_client_ref.append(self)
|
||
|
|
|
||
|
|
fake_profiles = _patch_profiles({
|
||
|
|
"gpt-5": "http://llm-gw/v1",
|
||
|
|
"Qwen3-Embedding-4B": "https://api.siliconflow.cn/v1",
|
||
|
|
})
|
||
|
|
|
||
|
|
import webapp.services.profile_manager as pm_mod
|
||
|
|
monkeypatch.setattr(pm_mod.profile_manager, "list_all", lambda: fake_profiles)
|
||
|
|
|
||
|
|
from ragas.llms.base import InstructorBaseRagasLLM
|
||
|
|
from ragas.embeddings.base import BaseRagasEmbedding
|
||
|
|
|
||
|
|
emb_factory_calls = []
|
||
|
|
|
||
|
|
def _fake_emb_factory(provider, model, client):
|
||
|
|
emb_factory_calls.append({"model": model, "client": client})
|
||
|
|
return MagicMock(spec=BaseRagasEmbedding)
|
||
|
|
|
||
|
|
call_count = [0]
|
||
|
|
|
||
|
|
def _client_factory(**kwargs):
|
||
|
|
call_count[0] += 1
|
||
|
|
if "llm-gw" in kwargs.get("base_url", ""):
|
||
|
|
return _LLMClient(**kwargs)
|
||
|
|
return _EmbClient(**kwargs)
|
||
|
|
|
||
|
|
with patch("rag_eval.metrics.factory.AsyncOpenAI", side_effect=_client_factory), \
|
||
|
|
patch("rag_eval.metrics.factory.llm_factory", return_value=MagicMock(spec=InstructorBaseRagasLLM)), \
|
||
|
|
patch("rag_eval.metrics.factory.embedding_factory", side_effect=_fake_emb_factory):
|
||
|
|
build_models("gpt-5", "Qwen3-Embedding-4B", _settings())
|
||
|
|
|
||
|
|
assert len(emb_factory_calls) == 1
|
||
|
|
# The client passed to embedding_factory must be the siliconflow client, not the LLM client.
|
||
|
|
assert isinstance(emb_factory_calls[0]["client"], _EmbClient), (
|
||
|
|
"embedding_factory should receive the embedding-model client, not the judge-model client"
|
||
|
|
)
|