77
88from __future__ import annotations
99
10- import copy
1110import json
1211from collections .abc import AsyncIterator , Iterator
1312from contextlib import asynccontextmanager
1918from pydantic import Field
2019
2120from verifiers .v1 .decorators import reward
22- from verifiers .v1 .dialects .responses import ResponsesDialect
21+ from verifiers .v1 .dialects .responses import ResponsesDialect , messages_to_wire
2322from verifiers .v1 .mcp import SharedToolsetConfig , Toolset
2423from verifiers .v1 .state import State
2524from verifiers .v1 .task import Task , TaskConfig , TaskData
@@ -44,7 +43,7 @@ class NeMoGymTaskConfig(TaskConfig):
4443
4544
4645class NeMoGymConfig (TasksetConfig ):
47- dataset_path : Path | None = None
46+ dataset_path : Path
4847 """JSONL rows containing ``responses_create_params`` and verifier metadata."""
4948
5049 task : NeMoGymTaskConfig = NeMoGymTaskConfig ()
@@ -77,7 +76,7 @@ class _NeMoGymToolset(Toolset[SharedToolsetConfig, NeMoGymState]):
7776 rollout session on every request.
7877 """
7978
80- TOOL_PREFIX = "nemo_gym"
79+ TOOL_PREFIX = None
8180
8281 def _register (self , mcp : FastMCP ) -> None :
8382 server = mcp ._mcp_server
@@ -124,8 +123,6 @@ async def list_tools(self) -> list[MCPTool]:
124123 async def call_tool (self , name : str , arguments : dict [str , Any ]) -> CallToolResult :
125124 from mcp .types import CallToolResult , TextContent
126125
127- if name not in self .state .tool_names :
128- raise ValueError (f"unknown NeMo Gym tool: { name } " )
129126 if self .state .mcp_url is not None :
130127 async with self ._upstream () as session :
131128 return await session .call_tool (name , arguments )
@@ -157,54 +154,25 @@ def _trace_to_nemo_response(
157154 f"NeMo Gym scoring requires exactly one trace branch, got { len (branches )} "
158155 )
159156
160- known_names = set (tool_names )
161- known_names .update (
157+ known_names = set (tool_names ) | {
162158 spec ["name" ]
163159 for spec in responses_create_params .get ("tools" ) or []
164160 if spec .get ("type" ) == "function" and isinstance (spec .get ("name" ), str )
165- )
166- aliases = {
167- name : {
168- name ,
169- f"nemo_gym_{ name } " ,
170- f"nemo_gym__{ name } " ,
171- f"mcp__nemo_gym__{ name } " ,
172- }
173- for name in known_names
174- }
175- response_item_types = {
176- "reasoning" ,
177- "message" ,
178- "function_call" ,
179- "mcp_call" ,
180- "mcp_list_tools" ,
181- "mcp_approval_request" ,
182161 }
162+ response_item_types = {"reasoning" , "message" , "function_call" }
183163 output : list [dict [str , Any ]] = []
184164 started = False
185165
186166 for node in branches [0 ].nodes :
187167 message = node .message
188168 if isinstance (message , AssistantMessage ) and node .sampled :
189169 started = True
190- provider_items = message .provider_state or []
191- if provider_items and all (
170+ if message .provider_state and not all (
192171 isinstance (item , dict ) and item .get ("type" ) in response_item_types
193- for item in provider_items
172+ for item in message . provider_state
194173 ):
195- items = copy .deepcopy (provider_items )
196- for item in items :
197- if item .get ("type" ) != "function_call" :
198- continue
199- emitted_name = str (item .get ("name" ))
200- for raw_name in sorted (known_names , key = len , reverse = True ):
201- if emitted_name in aliases [raw_name ]:
202- item ["name" ] = raw_name
203- break
204- output .extend (items )
205- continue
206-
207- if message .reasoning_content :
174+ message = message .model_copy (update = {"provider_state" : None })
175+ if message .reasoning_content and not message .provider_state :
208176 output .append (
209177 {
210178 "id" : f"rs_{ trace .id } _{ len (output )} " ,
@@ -217,49 +185,24 @@ def _trace_to_nemo_response(
217185 ],
218186 }
219187 )
220- for call in message .tool_calls or []:
221- name = call .name
222- for raw_name in sorted (known_names , key = len , reverse = True ):
223- if name in aliases [raw_name ]:
224- name = raw_name
225- break
226- output .append (
227- {
228- "id" : call .id ,
229- "type" : "function_call" ,
230- "call_id" : call .id ,
231- "name" : name ,
232- "arguments" : call .arguments ,
233- "status" : "completed" ,
234- }
235- )
236- if message .content :
237- output .append (
238- {
239- "id" : f"msg_{ trace .id } _{ len (output )} " ,
240- "type" : "message" ,
241- "role" : "assistant" ,
242- "status" : "completed" ,
243- "content" : [
244- {
245- "type" : "output_text" ,
246- "text" : message .content ,
247- "annotations" : [],
248- }
249- ],
250- }
251- )
188+ output .extend (dict (item ) for item in messages_to_wire ([message ]))
252189 elif started and isinstance (message , ToolMessage ):
253190 output .append (
254191 {
255- "id" : f"fco_{ trace .id } _{ len (output )} " ,
256192 "type" : "function_call_output" ,
257193 "call_id" : message .tool_call_id ,
258194 "output" : content_text (message .content ),
259- "status" : "completed" ,
260195 }
261196 )
262197
198+ for item in output :
199+ if item .get ("type" ) != "function_call" :
200+ continue
201+ name = str (item .get ("name" , "" ))
202+ bare_name = name .removeprefix ("_" )
203+ if bare_name in known_names :
204+ item ["name" ] = bare_name
205+
263206 model = responses_create_params .get ("model" ) or "verifiers"
264207 return {
265208 "id" : f"resp_{ trace .id } " ,
@@ -344,29 +287,15 @@ class NeMoGymTaskset(Taskset[NeMoGymTask, NeMoGymConfig]):
344287 tools = (_NeMoGymToolset ,)
345288
346289 def load (self ) -> Iterator [NeMoGymTask ]:
347- if self .config .dataset_path is None :
348- raise ValueError ("NeMoGymConfig.dataset_path is required" )
349290 path = self .config .dataset_path .expanduser ().resolve ()
350291 dialect = ResponsesDialect ()
351292 count = 0
352293
353294 with path .open (encoding = "utf-8" ) as dataset :
354- for line_number , line in enumerate ( dataset , start = 1 ) :
295+ for line in dataset :
355296 if not line .strip ():
356297 continue
357- try :
358- row = json .loads (line )
359- except json .JSONDecodeError as exc :
360- raise ValueError (
361- f"invalid JSON in { path } line { line_number } : { exc .msg } "
362- ) from exc
363- if not isinstance (row , dict ) or not isinstance (
364- row .get ("responses_create_params" ), dict
365- ):
366- raise ValueError (
367- f"{ path } line { line_number } must be an object with "
368- "responses_create_params"
369- )
298+ row = json .loads (line )
370299 params = row ["responses_create_params" ]
371300 prompt , _ = dialect .parse_request (params )
372301 yield NeMoGymTask (
0 commit comments