From f4dbc5af566d2f78265cc3e2f1863f740e756b1b Mon Sep 17 00:00:00 2001 From: Nitjsefnie Date: Wed, 29 Jul 2026 09:38:21 +0200 Subject: [PATCH 1/4] fix(pipeline): read LLM model through public model_name contract Add a public read-only LLMProvider.model_name property with a default implementation returning None so legacy providers keep working. Implement it on every concrete LLM provider in the repo to return the actual model identifier. Update Pipeline._run_metrics to use the public property while keeping the config fallback. Fixes #51 Co-Authored-By: Kimi K2.7 Code --- openagent_eval/core/pipeline.py | 2 +- openagent_eval/providers/base/llm.py | 16 +++ openagent_eval/providers/llm/anthropic.py | 5 + openagent_eval/providers/llm/gemini.py | 5 + openagent_eval/providers/llm/groq.py | 5 + openagent_eval/providers/llm/mock.py | 5 + openagent_eval/providers/llm/ollama.py | 5 + openagent_eval/providers/llm/openai.py | 5 + openagent_eval/providers/llm/openrouter.py | 5 + .../test_core/test_pipeline_model_name.py | 100 ++++++++++++++++++ 10 files changed, 152 insertions(+), 1 deletion(-) create mode 100644 tests/unit/test_core/test_pipeline_model_name.py diff --git a/openagent_eval/core/pipeline.py b/openagent_eval/core/pipeline.py index b5c11d7..993f35e 100644 --- a/openagent_eval/core/pipeline.py +++ b/openagent_eval/core/pipeline.py @@ -221,7 +221,7 @@ def _run_metrics( prompt_tokens = token_usage.prompt_tokens if token_usage else 0 completion_tokens = token_usage.completion_tokens if token_usage else 0 provider_name = getattr(self._llm, "name", None) - model_name = getattr(self._llm, "_model", None) or self.config.llm.model + model_name = getattr(self._llm, "model_name", None) or self.config.llm.model for name, metric in self._metrics: try: diff --git a/openagent_eval/providers/base/llm.py b/openagent_eval/providers/base/llm.py index eaae72b..d4f09b7 100644 --- a/openagent_eval/providers/base/llm.py +++ b/openagent_eval/providers/base/llm.py @@ -82,6 +82,22 @@ async def get_token_count(self, text: str) -> int: """ pass + @property + def model_name(self) -> str | None: + """Return the model identifier currently in use by this provider. + + Subclasses should override this property to report the actual model + identifier (e.g. ``"gpt-4o"``). The default implementation returns + ``None`` so that legacy third-party providers that do not define it + keep working; callers should fall back to ``config.llm.model`` when + this property returns ``None``. + + Returns: + The model identifier in use, or ``None`` if the provider does not + expose one. + """ + return None + def validate_inputs(self, **kwargs: Any) -> None: # noqa: B027 """Validate provider inputs before execution. diff --git a/openagent_eval/providers/llm/anthropic.py b/openagent_eval/providers/llm/anthropic.py index 1c28fc7..da4f3a1 100644 --- a/openagent_eval/providers/llm/anthropic.py +++ b/openagent_eval/providers/llm/anthropic.py @@ -139,6 +139,11 @@ def __init__( self._client = anthropic.AsyncAnthropic(api_key=api_key) + @property + def model_name(self) -> str | None: + """Return the configured Anthropic model identifier.""" + return self.model + async def generate(self, prompt: str, **kwargs: Any) -> str: """Generate a response from the LLM. diff --git a/openagent_eval/providers/llm/gemini.py b/openagent_eval/providers/llm/gemini.py index 06ab8c9..1ba4528 100644 --- a/openagent_eval/providers/llm/gemini.py +++ b/openagent_eval/providers/llm/gemini.py @@ -150,6 +150,11 @@ def __init__( original_error=exc, ) from exc + @property + def model_name(self) -> str | None: + """Return the configured Gemini model identifier.""" + return self._model + async def generate_with_usage(self, prompt: str, **kwargs: Any) -> LLMResponse: """Generate a response and return it with token usage and latency. diff --git a/openagent_eval/providers/llm/groq.py b/openagent_eval/providers/llm/groq.py index 728e8b0..b094f46 100644 --- a/openagent_eval/providers/llm/groq.py +++ b/openagent_eval/providers/llm/groq.py @@ -139,6 +139,11 @@ def __init__( original_error=e, ) from e + @property + def model_name(self) -> str | None: + """Return the configured Groq model identifier.""" + return self.model + async def generate(self, prompt: str, **kwargs: Any) -> str: """Generate a response from the LLM. diff --git a/openagent_eval/providers/llm/mock.py b/openagent_eval/providers/llm/mock.py index 61d33d0..0affc61 100644 --- a/openagent_eval/providers/llm/mock.py +++ b/openagent_eval/providers/llm/mock.py @@ -44,6 +44,11 @@ def __init__( self._model = getattr(config, "model", model) or model self._temperature = getattr(config, "temperature", temperature) or temperature + @property + def model_name(self) -> str | None: + """Return the configured mock model identifier.""" + return self._model + async def generate(self, prompt: str, **kwargs: Any) -> str: """Return a deterministic answer without calling any API. diff --git a/openagent_eval/providers/llm/ollama.py b/openagent_eval/providers/llm/ollama.py index b08d065..0b57fb6 100644 --- a/openagent_eval/providers/llm/ollama.py +++ b/openagent_eval/providers/llm/ollama.py @@ -145,6 +145,11 @@ def __init__( headers={"Content-Type": "application/json"}, ) + @property + def model_name(self) -> str | None: + """Return the configured Ollama model identifier.""" + return self._model + async def generate(self, prompt: str, **kwargs: Any) -> str: """Generate a response from the Ollama model. diff --git a/openagent_eval/providers/llm/openai.py b/openagent_eval/providers/llm/openai.py index 1e9ad7e..000201b 100644 --- a/openagent_eval/providers/llm/openai.py +++ b/openagent_eval/providers/llm/openai.py @@ -136,6 +136,11 @@ def __init__( self._client = AsyncOpenAI(api_key=self._api_key) self._encoding: tiktoken.Encoding | None = None + @property + def model_name(self) -> str | None: + """Return the configured OpenAI model identifier.""" + return self._model + def _get_encoding(self) -> tiktoken.Encoding: """Get or create tiktoken encoding for the configured model. diff --git a/openagent_eval/providers/llm/openrouter.py b/openagent_eval/providers/llm/openrouter.py index c7163f2..b23f053 100644 --- a/openagent_eval/providers/llm/openrouter.py +++ b/openagent_eval/providers/llm/openrouter.py @@ -122,6 +122,11 @@ def __init__( self.max_tokens = max_tokens self.base_url = base_url.rstrip("/") + @property + def model_name(self) -> str | None: + """Return the configured OpenRouter model identifier.""" + return self.model + async def generate_with_usage(self, prompt: str, **kwargs: Any) -> "LLMResponse": """Generate a response and return it with token usage and latency. diff --git a/tests/unit/test_core/test_pipeline_model_name.py b/tests/unit/test_core/test_pipeline_model_name.py new file mode 100644 index 0000000..697e805 --- /dev/null +++ b/tests/unit/test_core/test_pipeline_model_name.py @@ -0,0 +1,100 @@ +"""Regression test for issue #51. + +Verifies that ``openagent_eval/core/pipeline.py`` reads the LLM model +identifier through the public ``LLMProvider.model_name`` property rather than +reaching into a private ``_model`` attribute. +""" + +from __future__ import annotations + +from typing import Any + +import pytest + +from openagent_eval.config.models import ( + Config, + DatasetConfig, + LLMConfig, + MetricsConfig, + ReportConfig, + RetrieverConfig, +) +from openagent_eval.core.pipeline import Pipeline +from openagent_eval.metrics.base import BaseMetric, MetricResult +from openagent_eval.providers.base.llm import LLMProvider +from openagent_eval.providers.models import LLMResponse, TokenUsage + + +class _ModelNameCapturingMetric(BaseMetric): + """Metric that records the ``model`` argument it receives.""" + + name = "model_name_capture" + description = "Captures the model name passed by the pipeline" + + def __init__(self) -> None: + self.captured_model: str | None = None + + def evaluate(self, **kwargs: Any) -> MetricResult: + self.captured_model = kwargs.get("model") + return MetricResult(score=1.0, reason="captured") + + +class _FakeLLMProvider(LLMProvider): + """Fake LLM provider that exposes ``model_name`` but not ``_model``. + + This proves the pipeline uses the public property: the provider's + ``model_name`` differs from ``config.llm.model``, so a fallback to the + config value is visibly distinguishable from a real read. + """ + + name = "fake" + description = "Fake provider for regression testing" + + def __init__(self, model_name: str) -> None: + # Intentionally no ``_model`` attribute. + self._reported_model = model_name + + @property + def model_name(self) -> str | None: + return self._reported_model + + async def generate(self, prompt: str, **kwargs: Any) -> str: + return "fake answer" + + async def get_token_count(self, text: str) -> int: + return 0 + + async def generate_with_usage(self, prompt: str, **kwargs: Any) -> LLMResponse: + return LLMResponse( + content="fake answer", + model=self._reported_model, + usage=TokenUsage(prompt_tokens=1, completion_tokens=1, total_tokens=2), + provider=self.name, + latency_ms=1.0, + ) + + +@pytest.mark.asyncio +async def test_pipeline_uses_public_model_name_property() -> None: + """A provider exposing only ``model_name`` has it forwarded to metrics.""" + config_model = "config-model" + provider_model = "provider-specific-model" + + config = Config( + dataset=DatasetConfig(path="data/questions.json"), + llm=LLMConfig(provider="fake", model=config_model), + retriever=RetrieverConfig(provider="mock"), + metrics=MetricsConfig(), + report=ReportConfig(), + parallel=False, + ) + + llm = _FakeLLMProvider(model_name=provider_model) + capture_metric = _ModelNameCapturingMetric() + pipeline = Pipeline(config, llm=llm, metrics=[("model_name_capture", capture_metric)]) + + await pipeline.execute( + [{"question": "What is RAG?", "ground_truth": "RAG is retrieval augmented generation."}] + ) + + assert capture_metric.captured_model == provider_model From 73c2dc92365bdfaea2b1f8496501c90422f06fbd Mon Sep 17 00:00:00 2001 From: Nitjsefnie Date: Wed, 29 Jul 2026 09:39:00 +0200 Subject: [PATCH 2/4] fix(pipeline): pass ground_truth_contexts through public retriever contract Add ground_truth_contexts as an explicit keyword parameter to Retriever.retrieve with a None default. Update every concrete retriever in the repo to accept the parameter; retrievers that do not use it ignore it without changing behaviour. Remove the name == "mock" branch from Pipeline._retrieve and always forward the ground-truth contexts. Fixes #53 Co-Authored-By: Kimi K2.7 Code --- openagent_eval/core/pipeline.py | 9 +- openagent_eval/providers/base/retriever.py | 12 ++- openagent_eval/providers/retrievers/bm25.py | 8 +- openagent_eval/providers/retrievers/chroma.py | 8 +- .../providers/retrievers/elasticsearch.py | 8 +- openagent_eval/providers/retrievers/faiss.py | 8 +- openagent_eval/providers/retrievers/http.py | 8 +- openagent_eval/providers/retrievers/memory.py | 8 +- openagent_eval/providers/retrievers/mock.py | 14 +++- .../providers/retrievers/pgvector.py | 8 +- .../providers/retrievers/pinecone.py | 8 +- openagent_eval/providers/retrievers/qdrant.py | 8 +- .../providers/retrievers/weaviate.py | 8 +- .../test_pipeline_retriever_contract.py | 83 +++++++++++++++++++ 14 files changed, 177 insertions(+), 21 deletions(-) create mode 100644 tests/unit/test_core/test_pipeline_retriever_contract.py diff --git a/openagent_eval/core/pipeline.py b/openagent_eval/core/pipeline.py index 993f35e..d3d0e99 100644 --- a/openagent_eval/core/pipeline.py +++ b/openagent_eval/core/pipeline.py @@ -167,12 +167,9 @@ async def _retrieve( """Retrieve contexts for a question, or fall back to dataset context.""" if self._retriever is not None: try: - if getattr(self._retriever, "name", None) == "mock": - docs = await self._retriever.retrieve( - question, k=self._k, ground_truth_contexts=gt_contexts - ) - else: - docs = await self._retriever.retrieve(question, k=self._k) + docs = await self._retriever.retrieve( + question, k=self._k, ground_truth_contexts=gt_contexts + ) return [doc.content for doc in docs] except Exception: # Retrieval failure -> fall back to any dataset-provided context. diff --git a/openagent_eval/providers/base/retriever.py b/openagent_eval/providers/base/retriever.py index 7289933..0b48d73 100644 --- a/openagent_eval/providers/base/retriever.py +++ b/openagent_eval/providers/base/retriever.py @@ -43,12 +43,22 @@ async def retrieve(self, query: str, k: int = 5) -> list[Document]: description: str @abstractmethod - async def retrieve(self, query: str, k: int = 5) -> list[Document]: + async def retrieve( + self, + query: str, + k: int = 5, + *, + ground_truth_contexts: list[str] | None = None, + ) -> list[Document]: """Retrieve relevant documents for a given query. Args: query: The search query string. k: Number of documents to retrieve (default: 5). + ground_truth_contexts: Optional ground-truth contexts supplied by the + dataset. Retrievers that can use them (e.g. the mock retriever) + may return them; retrievers that do not use them must accept and + ignore the parameter. Returns: List of Document objects matching the query. diff --git a/openagent_eval/providers/retrievers/bm25.py b/openagent_eval/providers/retrievers/bm25.py index b755383..12faad9 100644 --- a/openagent_eval/providers/retrievers/bm25.py +++ b/openagent_eval/providers/retrievers/bm25.py @@ -108,7 +108,13 @@ async def _ensure_index(self) -> None: original_error=exc, ) from exc - async def retrieve(self, query: str, k: int = 5) -> list[Document]: + async def retrieve( + self, + query: str, + k: int = 5, + *, + ground_truth_contexts: list[str] | None = None, + ) -> list[Document]: """Retrieve the ``k`` highest BM25-scoring documents.""" self.validate_inputs(query=query, k=k) k = k or self._k diff --git a/openagent_eval/providers/retrievers/chroma.py b/openagent_eval/providers/retrievers/chroma.py index 1294e83..cc3cc45 100644 --- a/openagent_eval/providers/retrievers/chroma.py +++ b/openagent_eval/providers/retrievers/chroma.py @@ -125,7 +125,13 @@ def __init__( original_error=exc, ) from exc - async def retrieve(self, query: str, k: int = 5) -> list[Document]: + async def retrieve( + self, + query: str, + k: int = 5, + *, + ground_truth_contexts: list[str] | None = None, + ) -> list[Document]: """Retrieve relevant documents for a given query. Args: diff --git a/openagent_eval/providers/retrievers/elasticsearch.py b/openagent_eval/providers/retrievers/elasticsearch.py index 78346a5..2a12e1a 100644 --- a/openagent_eval/providers/retrievers/elasticsearch.py +++ b/openagent_eval/providers/retrievers/elasticsearch.py @@ -77,7 +77,13 @@ def __init__( original_error=exc, ) from exc - async def retrieve(self, query: str, k: int = 5) -> list[Document]: + async def retrieve( + self, + query: str, + k: int = 5, + *, + ground_truth_contexts: list[str] | None = None, + ) -> list[Document]: """Run a lexical or kNN search and normalize scores.""" self.validate_inputs(query=query, k=k) try: diff --git a/openagent_eval/providers/retrievers/faiss.py b/openagent_eval/providers/retrievers/faiss.py index f960007..c6fad6c 100644 --- a/openagent_eval/providers/retrievers/faiss.py +++ b/openagent_eval/providers/retrievers/faiss.py @@ -102,7 +102,13 @@ async def _ensure_index(self) -> None: original_error=exc, ) from exc - async def retrieve(self, query: str, k: int = 5) -> list[Document]: + async def retrieve( + self, + query: str, + k: int = 5, + *, + ground_truth_contexts: list[str] | None = None, + ) -> list[Document]: """Embed the query and search the FAISS index.""" self.validate_inputs(query=query, k=k) k = k or self._k diff --git a/openagent_eval/providers/retrievers/http.py b/openagent_eval/providers/retrievers/http.py index 30ed17a..0a64c51 100644 --- a/openagent_eval/providers/retrievers/http.py +++ b/openagent_eval/providers/retrievers/http.py @@ -121,7 +121,13 @@ def __init__( self._score_mode = score_mode self._timeout = timeout - async def retrieve(self, query: str, k: int = 5) -> list[Document]: + async def retrieve( + self, + query: str, + k: int = 5, + *, + ground_truth_contexts: list[str] | None = None, + ) -> list[Document]: """Send the query to the endpoint and map the response to Documents.""" self.validate_inputs(query=query, k=k) diff --git a/openagent_eval/providers/retrievers/memory.py b/openagent_eval/providers/retrievers/memory.py index be0b4b6..29d4c96 100644 --- a/openagent_eval/providers/retrievers/memory.py +++ b/openagent_eval/providers/retrievers/memory.py @@ -120,7 +120,13 @@ async def _ensure_index(self) -> None: original_error=exc, ) from exc - async def retrieve(self, query: str, k: int = 5) -> list[Document]: + async def retrieve( + self, + query: str, + k: int = 5, + *, + ground_truth_contexts: list[str] | None = None, + ) -> list[Document]: """Retrieve the ``k`` most similar documents to ``query``. Args: diff --git a/openagent_eval/providers/retrievers/mock.py b/openagent_eval/providers/retrievers/mock.py index d02654a..27c9c28 100644 --- a/openagent_eval/providers/retrievers/mock.py +++ b/openagent_eval/providers/retrievers/mock.py @@ -30,7 +30,14 @@ def __init__(self, collection_name: str = "mock", **_: Any) -> None: """ self.collection_name = collection_name - async def retrieve(self, query: str, k: int = 5, **kwargs: Any) -> list[Document]: + async def retrieve( + self, + query: str, + k: int = 5, + *, + ground_truth_contexts: list[str] | None = None, + **kwargs: Any, + ) -> list[Document]: """Return retrieved documents without a vector store. Args: @@ -44,8 +51,7 @@ async def retrieve(self, query: str, k: int = 5, **kwargs: Any) -> list[Document """ self.validate_inputs(query=query, k=k) - gt_contexts = kwargs.get("ground_truth_contexts") - if gt_contexts: + if ground_truth_contexts: return [ Document( content=ctx, @@ -53,7 +59,7 @@ async def retrieve(self, query: str, k: int = 5, **kwargs: Any) -> list[Document score=1.0, id=f"gt-{i}", ) - for i, ctx in enumerate(gt_contexts[:k]) + for i, ctx in enumerate(ground_truth_contexts[:k]) ] docs: list[Document] = [] diff --git a/openagent_eval/providers/retrievers/pgvector.py b/openagent_eval/providers/retrievers/pgvector.py index aaf7acf..fb8280b 100644 --- a/openagent_eval/providers/retrievers/pgvector.py +++ b/openagent_eval/providers/retrievers/pgvector.py @@ -88,7 +88,13 @@ async def _ensure_connection(self) -> None: original_error=exc, ) from exc - async def retrieve(self, query: str, k: int = 5) -> list[Document]: + async def retrieve( + self, + query: str, + k: int = 5, + *, + ground_truth_contexts: list[str] | None = None, + ) -> list[Document]: """Embed the query and run a similarity SQL query.""" self.validate_inputs(query=query, k=k) diff --git a/openagent_eval/providers/retrievers/pinecone.py b/openagent_eval/providers/retrievers/pinecone.py index cd95bde..41c1e41 100644 --- a/openagent_eval/providers/retrievers/pinecone.py +++ b/openagent_eval/providers/retrievers/pinecone.py @@ -69,7 +69,13 @@ def __init__( original_error=exc, ) from exc - async def retrieve(self, query: str, k: int = 5) -> list[Document]: + async def retrieve( + self, + query: str, + k: int = 5, + *, + ground_truth_contexts: list[str] | None = None, + ) -> list[Document]: """Embed the query and query the Pinecone index.""" self.validate_inputs(query=query, k=k) try: diff --git a/openagent_eval/providers/retrievers/qdrant.py b/openagent_eval/providers/retrievers/qdrant.py index 952ecbc..a43f2cb 100644 --- a/openagent_eval/providers/retrievers/qdrant.py +++ b/openagent_eval/providers/retrievers/qdrant.py @@ -80,7 +80,13 @@ def __init__( original_error=exc, ) from exc - async def retrieve(self, query: str, k: int = 5) -> list[Document]: + async def retrieve( + self, + query: str, + k: int = 5, + *, + ground_truth_contexts: list[str] | None = None, + ) -> list[Document]: """Embed the query and search the Qdrant collection.""" self.validate_inputs(query=query, k=k) try: diff --git a/openagent_eval/providers/retrievers/weaviate.py b/openagent_eval/providers/retrievers/weaviate.py index 6a24614..e4ff946 100644 --- a/openagent_eval/providers/retrievers/weaviate.py +++ b/openagent_eval/providers/retrievers/weaviate.py @@ -73,7 +73,13 @@ def __init__( original_error=exc, ) from exc - async def retrieve(self, query: str, k: int = 5) -> list[Document]: + async def retrieve( + self, + query: str, + k: int = 5, + *, + ground_truth_contexts: list[str] | None = None, + ) -> list[Document]: """Search the Weaviate collection near the query text.""" self.validate_inputs(query=query, k=k) try: diff --git a/tests/unit/test_core/test_pipeline_retriever_contract.py b/tests/unit/test_core/test_pipeline_retriever_contract.py new file mode 100644 index 0000000..fb2a52b --- /dev/null +++ b/tests/unit/test_core/test_pipeline_retriever_contract.py @@ -0,0 +1,83 @@ +"""Regression test for issue #53. + +Verifies that ``openagent_eval/core/pipeline.py`` passes +``ground_truth_contexts`` to every retriever through the public +``Retriever.retrieve`` signature instead of special-casing the mock retriever +by name. +""" + +from __future__ import annotations + +import pytest + +from openagent_eval.config.models import ( + Config, + DatasetConfig, + LLMConfig, + MetricsConfig, + ReportConfig, + RetrieverConfig, +) +from openagent_eval.core.pipeline import Pipeline +from openagent_eval.providers.base.retriever import Retriever +from openagent_eval.providers.models import Document + + +class _NonMockRetriever(Retriever): + """Non-mock retriever that records the ground_truth_contexts it receives.""" + + name = "not-mock" + description = "Records ground_truth_contexts for regression testing" + + def __init__(self) -> None: + self.captured_ground_truth_contexts: list[str] | None = None + + async def retrieve( + self, + query: str, + k: int = 5, + *, + ground_truth_contexts: list[str] | None = None, + ) -> list[Document]: + self.captured_ground_truth_contexts = ground_truth_contexts + return [ + Document( + content="retrieved context", + score=1.0, + id="doc-1", + ) + ] + + +@pytest.mark.asyncio +async def test_pipeline_passes_ground_truth_contexts_to_non_mock_retriever() -> None: + """A retriever whose name is not ``mock`` still receives ground_truth_contexts.""" + expected_contexts = ["ground truth context one", "ground truth context two"] + + # A minimal fake LLM so the pipeline can run generation. The provider is not + # the subject of this test. + from openagent_eval.providers.llm.mock import MockLLMProvider + + config = Config( + dataset=DatasetConfig(path="data/questions.json"), + llm=LLMConfig(provider="mock", model="mock-model"), + retriever=RetrieverConfig(provider="not-mock"), + metrics=MetricsConfig(), + report=ReportConfig(), + parallel=False, + ) + + retriever = _NonMockRetriever() + pipeline = Pipeline(config, retriever=retriever, llm=MockLLMProvider()) + + await pipeline.execute( + [ + { + "question": "What is RAG?", + "ground_truth": "RAG is retrieval augmented generation.", + "ground_truth_contexts": expected_contexts, + } + ] + ) + + assert retriever.captured_ground_truth_contexts == expected_contexts From 0bed727eb535a29fbfa713036d3276150f7716f1 Mon Sep 17 00:00:00 2001 From: Nitjsefnie Date: Wed, 29 Jul 2026 09:43:06 +0200 Subject: [PATCH 3/4] style(tests): format regression test for issue #51 Co-Authored-By: Kimi K2.7 Code --- tests/unit/test_core/test_pipeline_model_name.py | 11 +++++++++-- 1 file changed, 9 insertions(+), 2 deletions(-) diff --git a/tests/unit/test_core/test_pipeline_model_name.py b/tests/unit/test_core/test_pipeline_model_name.py index 697e805..e800d10 100644 --- a/tests/unit/test_core/test_pipeline_model_name.py +++ b/tests/unit/test_core/test_pipeline_model_name.py @@ -91,10 +91,17 @@ async def test_pipeline_uses_public_model_name_property() -> None: llm = _FakeLLMProvider(model_name=provider_model) capture_metric = _ModelNameCapturingMetric() - pipeline = Pipeline(config, llm=llm, metrics=[("model_name_capture", capture_metric)]) + pipeline = Pipeline( + config, llm=llm, metrics=[("model_name_capture", capture_metric)] + ) await pipeline.execute( - [{"question": "What is RAG?", "ground_truth": "RAG is retrieval augmented generation."}] + [ + { + "question": "What is RAG?", + "ground_truth": "RAG is retrieval augmented generation.", + } + ] ) assert capture_metric.captured_model == provider_model From df6ca083e1f358543b15f9ce20a127d31fd2b73f Mon Sep 17 00:00:00 2001 From: Nitjsefnie Date: Wed, 29 Jul 2026 10:30:05 +0200 Subject: [PATCH 4/4] fix(pipeline): detect ground_truth_contexts support for backward compatibility Commit 73c2dc9 unconditionally passed ground_truth_contexts to every retriever. Third-party retrievers written against the legacy retrieve(query, k=5) signature therefore raised TypeError, which the pre-existing bare except Exception in _retrieve swallowed silently, causing retrieval to degrade to the dataset fallback on every call. Detect capability once per retriever using inspect.signature on the bound retrieve method, cache the result on the pipeline instance, and pass the keyword only when the retriever supports it (either explicitly or via **kwargs). The legacy two-argument form is called otherwise. The name == "mock" check remains removed. Add a regression test proving a legacy-signature retriever's documents reach the pipeline result. Co-Authored-By: Kimi K2.7 Code --- openagent_eval/core/pipeline.py | 70 ++++++++++++++++--- .../test_pipeline_retriever_contract.py | 47 +++++++++++++ 2 files changed, 107 insertions(+), 10 deletions(-) diff --git a/openagent_eval/core/pipeline.py b/openagent_eval/core/pipeline.py index d3d0e99..ac0e6c1 100644 --- a/openagent_eval/core/pipeline.py +++ b/openagent_eval/core/pipeline.py @@ -11,6 +11,7 @@ from __future__ import annotations +import inspect from dataclasses import dataclass, field from typing import Any @@ -65,6 +66,7 @@ def __init__( self._metrics: list[tuple[str, BaseMetric]] = metrics or [] self._executor = executor self._k = self._resolve_k() + self._retriever_supports_ground_truth_contexts: bool | None = None # ------------------------------------------------------------------ # # Public API # @@ -121,8 +123,13 @@ async def _evaluate_item( # 3. Metrics metrics, metric_errors = self._run_metrics( - question, answer, ground_truth, contexts, gt_contexts, - latency_ms, token_usage, + question, + answer, + ground_truth, + contexts, + gt_contexts, + latency_ms, + token_usage, ) return EvaluationResult( @@ -134,7 +141,9 @@ async def _evaluate_item( metadata={ "latency_ms": latency_ms, "prompt_tokens": token_usage.prompt_tokens if token_usage else None, - "completion_tokens": token_usage.completion_tokens if token_usage else None, + "completion_tokens": token_usage.completion_tokens + if token_usage + else None, "total_tokens": token_usage.total_tokens if token_usage else None, "metric_errors": metric_errors, **item.get("metadata", {}), @@ -161,15 +170,56 @@ async def _evaluate_item( # Steps # # ------------------------------------------------------------------ # + def _supports_ground_truth_contexts(self) -> bool: + """Return whether the injected retriever accepts ``ground_truth_contexts``. + + The result is cached on the pipeline instance because it is computed + with reflection and must not run inside the per-item retrieval hot loop. + A retriever is treated as supporting the parameter if its bound + ``retrieve`` method has a parameter named ``ground_truth_contexts`` or + accepts arbitrary keyword arguments (``**kwargs``). + """ + if self._retriever_supports_ground_truth_contexts is not None: + return self._retriever_supports_ground_truth_contexts + + if self._retriever is None: + self._retriever_supports_ground_truth_contexts = False + return False + + retrieve_method = getattr(self._retriever, "retrieve", None) + if retrieve_method is None: + self._retriever_supports_ground_truth_contexts = False + return False + + try: + sig = inspect.signature(retrieve_method) + except (ValueError, TypeError): + self._retriever_supports_ground_truth_contexts = False + return False + + for param in sig.parameters.values(): + if param.kind == inspect.Parameter.VAR_KEYWORD: + self._retriever_supports_ground_truth_contexts = True + return True + if param.name == "ground_truth_contexts": + self._retriever_supports_ground_truth_contexts = True + return True + + self._retriever_supports_ground_truth_contexts = False + return False + async def _retrieve( self, question: str, context: str | None, gt_contexts: list[str] ) -> list[str]: """Retrieve contexts for a question, or fall back to dataset context.""" if self._retriever is not None: try: - docs = await self._retriever.retrieve( - question, k=self._k, ground_truth_contexts=gt_contexts - ) + if self._supports_ground_truth_contexts(): + docs = await self._retriever.retrieve( + question, k=self._k, ground_truth_contexts=gt_contexts + ) + else: + docs = await self._retriever.retrieve(question, k=self._k) return [doc.content for doc in docs] except Exception: # Retrieval failure -> fall back to any dataset-provided context. @@ -190,7 +240,9 @@ async def _generate( return "", None, None try: - response = await self._llm.generate_with_usage(prompt, ground_truth=ground_truth) + response = await self._llm.generate_with_usage( + prompt, ground_truth=ground_truth + ) return response.content, response.usage, response.latency_ms except Exception: # Generation failure -> empty answer; metrics will report accordingly. @@ -290,9 +342,7 @@ def _compute_summary( if metric_counts[name] > 0 } - total_tokens = sum( - (r.metadata.get("total_tokens") or 0) for r in results - ) + total_tokens = sum((r.metadata.get("total_tokens") or 0) for r in results) latencies = [ r.metadata.get("latency_ms") for r in results diff --git a/tests/unit/test_core/test_pipeline_retriever_contract.py b/tests/unit/test_core/test_pipeline_retriever_contract.py index fb2a52b..3e77bde 100644 --- a/tests/unit/test_core/test_pipeline_retriever_contract.py +++ b/tests/unit/test_core/test_pipeline_retriever_contract.py @@ -49,6 +49,16 @@ async def retrieve( ] +class _LegacyRetriever(Retriever): + """Out-of-tree retriever written against the old two-argument signature.""" + + name = "legacy" + description = "Legacy retriever with no ground_truth_contexts support" + + async def retrieve(self, query: str, k: int = 5) -> list[Document]: + return [Document(content=f"legacy doc for {query}", score=1.0, id="legacy-1")] + + @pytest.mark.asyncio async def test_pipeline_passes_ground_truth_contexts_to_non_mock_retriever() -> None: """A retriever whose name is not ``mock`` still receives ground_truth_contexts.""" @@ -81,3 +91,40 @@ async def test_pipeline_passes_ground_truth_contexts_to_non_mock_retriever() -> ) assert retriever.captured_ground_truth_contexts == expected_contexts + + +@pytest.mark.asyncio +async def test_pipeline_keeps_legacy_retriever_working() -> None: + """A retriever with the old two-argument signature still returns documents. + + Before the capability-detection fix, the pipeline unconditionally passed + ``ground_truth_contexts`` to every retriever. A legacy retriever raised + ``TypeError``, which was swallowed by the bare ``except Exception`` in + ``_retrieve``, causing it to silently return no documents. This test + asserts that the legacy retriever's documents actually reach the result. + """ + from openagent_eval.providers.llm.mock import MockLLMProvider + + config = Config( + dataset=DatasetConfig(path="data/questions.json"), + llm=LLMConfig(provider="mock", model="mock-model"), + retriever=RetrieverConfig(provider="legacy"), + metrics=MetricsConfig(), + report=ReportConfig(), + parallel=False, + ) + + retriever = _LegacyRetriever() + pipeline = Pipeline(config, retriever=retriever, llm=MockLLMProvider()) + + result = await pipeline.execute( + [ + { + "question": "What is RAG?", + "ground_truth": "RAG is retrieval augmented generation.", + "ground_truth_contexts": ["expected gt context"], + } + ] + ) + + assert result.results[0].contexts == ["legacy doc for What is RAG?"]