11import asyncio
22import logging
33import types
4- from typing import List
4+ from typing import List , Literal
55
6- from attr import dataclass
76from openai .types .chat .chat_completion_assistant_message_param import ChatCompletionAssistantMessageParam
87
98from eval_protocol .models import EvaluationRow , Message
1817from pydantic_ai .messages import ModelMessage
1918from pydantic_ai ._utils import generate_tool_call_id
2019from pydantic_ai import Agent
20+ from pydantic_ai .usage import UsageLimits
2121from pydantic_ai .messages import (
2222 ModelRequest ,
2323 SystemPromptPart ,
2424 ToolReturnPart ,
2525 UserPromptPart ,
2626)
2727from pydantic_ai .providers .openai import OpenAIProvider
28- from typing_extensions import TypedDict
28+ from typing_extensions import Callable
2929
3030logger = 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 } " )
0 commit comments