Extract shared build_metric_registry factory (DRY)
Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
This commit is contained in:
@@ -0,0 +1,141 @@
|
||||
"""Tests for the optimization advisor's direct-LLM analyzer.
|
||||
|
||||
These tests inject a fake async chat client so no network call is made. They
|
||||
verify that analyze() uses a plain chat.completions call (not the removed
|
||||
langchain path), returns the text from choices[0].message.content, embeds the
|
||||
worked-example instructions and the low-sample contexts in the prompt, and
|
||||
selects the correct token parameter for reasoning vs. legacy models.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
|
||||
from rag_eval.advisor.llm_analyzer import analyze, _is_reasoning_model
|
||||
from rag_eval.advisor.rules import Diagnosis
|
||||
|
||||
|
||||
class _FakeMessage:
|
||||
def __init__(self, content: str) -> None:
|
||||
self.content = content
|
||||
|
||||
|
||||
class _FakeChoice:
|
||||
def __init__(self, content: str) -> None:
|
||||
self.message = _FakeMessage(content)
|
||||
|
||||
|
||||
class _FakeResponse:
|
||||
def __init__(self, content: str) -> None:
|
||||
self.choices = [_FakeChoice(content)]
|
||||
|
||||
|
||||
class _FakeCompletions:
|
||||
def __init__(self, captured: dict) -> None:
|
||||
self._captured = captured
|
||||
|
||||
async def create(self, **kwargs):
|
||||
self._captured.update(kwargs)
|
||||
return _FakeResponse("## faithfulness [警告]\n\n针对该问题的具体优化建议")
|
||||
|
||||
|
||||
class _FakeChat:
|
||||
def __init__(self, captured: dict) -> None:
|
||||
self.completions = _FakeCompletions(captured)
|
||||
|
||||
|
||||
class _FakeClient:
|
||||
def __init__(self, captured: dict) -> None:
|
||||
self.chat = _FakeChat(captured)
|
||||
self.closed = False
|
||||
|
||||
async def close(self) -> None:
|
||||
self.closed = True
|
||||
|
||||
|
||||
class _Settings:
|
||||
ragas_llm_max_tokens = 4096
|
||||
|
||||
|
||||
def _diagnosis() -> Diagnosis:
|
||||
return Diagnosis(
|
||||
metric="faithfulness",
|
||||
mean_score=0.55,
|
||||
threshold=0.7,
|
||||
severity="warning",
|
||||
root_causes=["生成未严格 grounding"],
|
||||
suggested_actions=["强化 grounding 约束"],
|
||||
low_samples=[
|
||||
{
|
||||
"sample_id": "s1",
|
||||
"question": "球管寿命如何评估?",
|
||||
"answer": "球管寿命约 3 年。",
|
||||
"ground_truth": "球管寿命取决于使用强度。",
|
||||
"contexts": "球管寿命与扫描负载相关 |||| 高负载会缩短寿命",
|
||||
"faithfulness": 0.4,
|
||||
}
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
def test_analyze_uses_direct_chat_and_returns_content() -> None:
|
||||
captured: dict = {}
|
||||
text = asyncio.run(
|
||||
analyze([_diagnosis()], "scn", "gpt-4o", _Settings(), chat_client=_FakeClient(captured))
|
||||
)
|
||||
assert "优化建议" in text
|
||||
assert captured["model"] == "gpt-4o"
|
||||
prompt = captured["messages"][0]["content"]
|
||||
assert "举例拆解" in prompt # worked-example instruction present
|
||||
assert "球管寿命与扫描负载相关" in prompt # low-sample contexts embedded
|
||||
assert "max_tokens" in captured # legacy model uses max_tokens
|
||||
assert "max_completion_tokens" not in captured
|
||||
|
||||
|
||||
def test_analyze_reasoning_model_uses_max_completion_tokens() -> None:
|
||||
captured: dict = {}
|
||||
asyncio.run(
|
||||
analyze([_diagnosis()], "scn", "gpt-5", _Settings(), chat_client=_FakeClient(captured))
|
||||
)
|
||||
assert "max_completion_tokens" in captured
|
||||
assert "max_tokens" not in captured
|
||||
|
||||
|
||||
def test_analyze_empty_diagnoses_returns_empty() -> None:
|
||||
assert asyncio.run(analyze([], "scn", "gpt-4o", _Settings())) == ""
|
||||
|
||||
|
||||
def test_analyze_closes_client_it_creates(monkeypatch) -> None:
|
||||
"""A self-created client is closed in-loop to avoid 'Event loop is closed'."""
|
||||
captured: dict = {}
|
||||
fake = _FakeClient(captured)
|
||||
|
||||
import openai
|
||||
import rag_eval.metrics.factory as factory_mod
|
||||
|
||||
monkeypatch.setattr(openai, "AsyncOpenAI", lambda **kwargs: fake)
|
||||
monkeypatch.setattr(
|
||||
factory_mod, "resolve_openai_client_kwargs", lambda *a, **k: {"api_key": "x"}
|
||||
)
|
||||
|
||||
# No chat_client passed → analyze() builds (and must close) its own client.
|
||||
text = asyncio.run(analyze([_diagnosis()], "scn", "gpt-4o", _Settings()))
|
||||
|
||||
assert "优化建议" in text
|
||||
assert fake.closed is True
|
||||
|
||||
|
||||
def test_analyze_does_not_close_injected_client() -> None:
|
||||
"""An injected client is owned by the caller and must not be closed."""
|
||||
fake = _FakeClient({})
|
||||
asyncio.run(analyze([_diagnosis()], "scn", "gpt-4o", _Settings(), chat_client=fake))
|
||||
assert fake.closed is False
|
||||
|
||||
|
||||
def test_is_reasoning_model_detection() -> None:
|
||||
assert _is_reasoning_model("gpt-5")
|
||||
assert _is_reasoning_model("gpt-5.5")
|
||||
assert _is_reasoning_model("o1-mini")
|
||||
assert _is_reasoning_model("o3")
|
||||
assert not _is_reasoning_model("gpt-4o")
|
||||
assert not _is_reasoning_model("deepseek-v4-flash")
|
||||
Reference in New Issue
Block a user