Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
65 changes: 56 additions & 9 deletions openagent_eval/core/pipeline.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@

from __future__ import annotations

import inspect
from dataclasses import dataclass, field
from typing import Any

Expand Down Expand Up @@ -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 #
Expand Down Expand Up @@ -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(
Expand All @@ -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", {}),
Expand All @@ -161,13 +170,51 @@ 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:
if getattr(self._retriever, "name", None) == "mock":
if self._supports_ground_truth_contexts():
docs = await self._retriever.retrieve(
question, k=self._k, ground_truth_contexts=gt_contexts
)
Expand All @@ -193,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.
Expand Down Expand Up @@ -221,7 +270,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:
Expand Down Expand Up @@ -293,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
Expand Down
16 changes: 16 additions & 0 deletions openagent_eval/providers/base/llm.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand Down
12 changes: 11 additions & 1 deletion openagent_eval/providers/base/retriever.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
5 changes: 5 additions & 0 deletions openagent_eval/providers/llm/anthropic.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand Down
5 changes: 5 additions & 0 deletions openagent_eval/providers/llm/gemini.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand Down
5 changes: 5 additions & 0 deletions openagent_eval/providers/llm/groq.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand Down
5 changes: 5 additions & 0 deletions openagent_eval/providers/llm/mock.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand Down
5 changes: 5 additions & 0 deletions openagent_eval/providers/llm/ollama.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand Down
5 changes: 5 additions & 0 deletions openagent_eval/providers/llm/openai.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand Down
5 changes: 5 additions & 0 deletions openagent_eval/providers/llm/openrouter.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand Down
8 changes: 7 additions & 1 deletion openagent_eval/providers/retrievers/bm25.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
8 changes: 7 additions & 1 deletion openagent_eval/providers/retrievers/chroma.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
8 changes: 7 additions & 1 deletion openagent_eval/providers/retrievers/elasticsearch.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
8 changes: 7 additions & 1 deletion openagent_eval/providers/retrievers/faiss.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
8 changes: 7 additions & 1 deletion openagent_eval/providers/retrievers/http.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand Down
8 changes: 7 additions & 1 deletion openagent_eval/providers/retrievers/memory.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
14 changes: 10 additions & 4 deletions openagent_eval/providers/retrievers/mock.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand All @@ -44,16 +51,15 @@ 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,
metadata={"mock": True, "source": "ground_truth_contexts"},
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] = []
Expand Down
8 changes: 7 additions & 1 deletion openagent_eval/providers/retrievers/pgvector.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand Down
Loading
Loading