Skip to content

Commit f17169e

Browse files
author
Dylan Huang
committed
TODO: refactor rolloutprocessor to not use __call__
1 parent 7239134 commit f17169e

3 files changed

Lines changed: 52 additions & 41 deletions

File tree

eval_protocol/pytest/default_pydantic_ai_rollout_processor.py

Lines changed: 48 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -1,9 +1,8 @@
11
import asyncio
22
import logging
33
import types
4-
from typing import List
4+
from typing import List, Literal
55

6-
from attr import dataclass
76
from openai.types.chat.chat_completion_assistant_message_param import ChatCompletionAssistantMessageParam
87

98
from eval_protocol.models import EvaluationRow, Message
@@ -18,14 +17,15 @@
1817
from pydantic_ai.messages import ModelMessage
1918
from pydantic_ai._utils import generate_tool_call_id
2019
from pydantic_ai import Agent
20+
from pydantic_ai.usage import UsageLimits
2121
from pydantic_ai.messages import (
2222
ModelRequest,
2323
SystemPromptPart,
2424
ToolReturnPart,
2525
UserPromptPart,
2626
)
2727
from pydantic_ai.providers.openai import OpenAIProvider
28-
from typing_extensions import TypedDict
28+
from typing_extensions import Callable
2929

3030
logger = logging.getLogger(__name__)
3131

@@ -34,9 +34,33 @@ class PydanticAgentRolloutProcessor(RolloutProcessor):
3434
"""Rollout processor for Pydantic AI agents. Mainly converts
3535
EvaluationRow.messages to and from Pydantic AI ModelMessage format."""
3636

37-
def __init__(self):
37+
def __init__(self, setup_agent: Callable[..., Agent], usage_limits: UsageLimits = None):
3838
# dummy model used for its helper functions for processing messages
3939
self.util = OpenAIModel("dummy-model", provider=OpenAIProvider(api_key="dummy"))
40+
self.setup_agent = setup_agent
41+
self.usage_limits = usage_limits
42+
43+
def _map_litellm_to_pydantic_ai(
44+
self, model_name: str
45+
) -> Literal[
46+
"openai",
47+
"deepseek",
48+
"azure",
49+
"openrouter",
50+
"grok",
51+
"fireworks",
52+
"together",
53+
]:
54+
mapping = {
55+
"fireworks_ai": "fireworks",
56+
"together_ai": "together",
57+
"xai": "grok",
58+
"azure_ai": "azure",
59+
}
60+
provider = model_name.split("/")[0]
61+
if provider in mapping:
62+
provider = mapping[provider]
63+
return provider # type: ignore
4064

4165
def __call__(self, rows: List[EvaluationRow], config: RolloutProcessorConfig) -> List[asyncio.Task[EvaluationRow]]:
4266
"""Create agent rollout tasks and return them for external handling."""
@@ -60,20 +84,28 @@ def __call__(self, rows: List[EvaluationRow], config: RolloutProcessorConfig) ->
6084
raise ValueError(
6185
"completion_params['model'] must be a dict mapping agent argument names to model config dicts (with 'model' and 'provider' keys)"
6286
)
63-
kwargs = {}
64-
for k, v in config.completion_params["model"].items():
65-
if v["model"] and v["model"].startswith("anthropic:"):
66-
kwargs[k] = AnthropicModel(
67-
v["model"].removeprefix("anthropic:"),
87+
kwargs: dict[str, OpenAIModel | GoogleModel | AnthropicModel] = {}
88+
for agent, model_config in config.completion_params["model"].items():
89+
if "model" not in model_config:
90+
raise ValueError(f"model_config for agent {agent} must contain a 'model' key")
91+
model_name = model_config["model"]
92+
if model_name.startswith("anthropic/"):
93+
kwargs[agent] = AnthropicModel(
94+
model_name.removeprefix("anthropic/"),
95+
)
96+
elif model_name.startswith("google/"):
97+
kwargs[agent] = GoogleModel(
98+
model_name.removeprefix("google/"),
6899
)
69-
elif v["model"] and v["model"].startswith("google:"):
70-
kwargs[k] = GoogleModel(
71-
v["model"].removeprefix("google:"),
100+
elif model_name.startswith("gemini/"):
101+
kwargs[agent] = GoogleModel(
102+
model_name.removeprefix("gemini/"),
72103
)
73104
else:
74-
kwargs[k] = OpenAIModel(
75-
v["model"],
76-
provider=v["provider"],
105+
provider = self._map_litellm_to_pydantic_ai(model_name)
106+
kwargs[agent] = OpenAIModel(
107+
model_name.removeprefix(f"{provider}/"),
108+
provider=provider,
77109
)
78110
agent = setup_agent(**kwargs)
79111
model = None
@@ -144,5 +176,4 @@ def convert_ep_message_to_pyd_message(self, message: Message, row: EvaluationRow
144176
)
145177
]
146178
)
147-
else:
148-
raise ValueError(f"Unknown role: {message.role}")
179+
raise ValueError(f"Unknown role: {message.role}")

eval_protocol/pytest/evaluation_test.py

Lines changed: 0 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -544,11 +544,6 @@ def _log_eval_error(status: Status, rows: Optional[List[EvaluationRow]] | None,
544544
row.input_metadata.row_id = generate_id(seed=0, index=index)
545545

546546
completion_params = kwargs["completion_params"]
547-
if completion_params and ("model" not in completion_params or not completion_params["model"]):
548-
raise ValueError(
549-
"No model provided. Please provide a model in the completion parameters object."
550-
)
551-
552547
# Create eval metadata with test function info and current commit hash
553548
eval_metadata = EvalMetadata(
554549
name=test_func.__name__,

tests/pytest/test_pydantic_multi_agent.py

Lines changed: 4 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -47,12 +47,8 @@ async def joke_factory(ctx: RunContext[None], count: int) -> list[str]:
4747

4848
@pytest.mark.asyncio
4949
@evaluation_test(
50-
input_messages=[Message(role="user", content="Tell me a joke.")],
50+
input_messages=[[Message(role="user", content="Tell me a joke.")]],
5151
completion_params=[
52-
# single agent
53-
{
54-
"model": "fireworks_ai/accounts/fireworks/models/kimi-k2-instruct",
55-
},
5652
# multi-agent
5753
{
5854
"joke_generation_model": {
@@ -62,21 +58,10 @@ async def joke_factory(ctx: RunContext[None], count: int) -> list[str]:
6258
"model": "fireworks_ai/accounts/fireworks/models/deepseek-v3p1",
6359
},
6460
},
65-
{
66-
"joke_generation_model": {
67-
"model": "fireworks_ai/accounts/fireworks/models/kimi-k2-instruct",
68-
},
69-
"joke_selection_model": {
70-
"model": "fireworks_ai/accounts/fireworks/models/kimi-k2-instruct",
71-
},
72-
},
7361
],
74-
rollout_processor=PydanticAgentRolloutProcessor(),
75-
rollout_processor_kwargs={
76-
"agent": setup_agent,
77-
# PydanticAgentRolloutProcessor will pass usage_limits into the "run" call
78-
"usage_limits": UsageLimits(request_limit=5, total_tokens_limit=1000),
79-
},
62+
rollout_processor=PydanticAgentRolloutProcessor.__init__(
63+
setup_agent, UsageLimits(request_limit=5, total_tokens_limit=1000)
64+
),
8065
mode="pointwise",
8166
)
8267
async def test_pydantic_multi_agent(row: EvaluationRow) -> EvaluationRow:

0 commit comments

Comments
 (0)