jx-react-agent 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- jx_react_agent-0.1.0.dist-info/METADATA +231 -0
- jx_react_agent-0.1.0.dist-info/RECORD +13 -0
- jx_react_agent-0.1.0.dist-info/WHEEL +5 -0
- jx_react_agent-0.1.0.dist-info/top_level.txt +1 -0
- jx_react_agent_autoload.pth +1 -0
- rewardkit_react/__init__.py +13 -0
- rewardkit_react/agent.py +537 -0
- rewardkit_react/registry.py +152 -0
- rewardkit_react/tool.py +86 -0
- rewardkit_react/tools/__init__.py +14 -0
- rewardkit_react/tools/documents.py +436 -0
- rewardkit_react/tools/files.py +325 -0
- rewardkit_react/tools/shell.py +80 -0
rewardkit_react/agent.py
ADDED
|
@@ -0,0 +1,537 @@
|
|
|
1
|
+
"""A small native-tool-calling ReAct backend for RewardKit 0.2.x."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import asyncio
|
|
6
|
+
import base64
|
|
7
|
+
import copy
|
|
8
|
+
import io
|
|
9
|
+
import json
|
|
10
|
+
import mimetypes
|
|
11
|
+
import os
|
|
12
|
+
import tempfile
|
|
13
|
+
import time
|
|
14
|
+
import uuid
|
|
15
|
+
from collections import Counter
|
|
16
|
+
from pathlib import Path
|
|
17
|
+
from typing import Any
|
|
18
|
+
|
|
19
|
+
import litellm
|
|
20
|
+
from rewardkit import AgentAttempt, AgentBackend
|
|
21
|
+
|
|
22
|
+
from rewardkit_react.registry import ToolRegistry, load_registry
|
|
23
|
+
from rewardkit_react.tool import ToolContext, ToolResult
|
|
24
|
+
|
|
25
|
+
_DEFAULT_MAX_STEPS = 20
|
|
26
|
+
_DEFAULT_MAX_TOOL_CHARS = 40_000
|
|
27
|
+
_DEFAULT_MAX_IMAGE_BYTES = 8 * 1024 * 1024
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
class ReActBackend(AgentBackend):
|
|
31
|
+
"""Explore a workspace through model tool calls, then submit a score."""
|
|
32
|
+
|
|
33
|
+
name = "react"
|
|
34
|
+
|
|
35
|
+
def __init__(self, judge: Any, cwd: str | None) -> None:
|
|
36
|
+
super().__init__(judge, cwd)
|
|
37
|
+
self.registry, self.load_warnings = load_registry()
|
|
38
|
+
|
|
39
|
+
async def run(
|
|
40
|
+
self, prompt: str, schema: dict[str, Any], label: str
|
|
41
|
+
) -> AgentAttempt:
|
|
42
|
+
model = self.judge.model or os.environ.get("REACT_JUDGE_MODEL")
|
|
43
|
+
if not model:
|
|
44
|
+
return AgentAttempt(
|
|
45
|
+
error=ValueError(
|
|
46
|
+
"the react judge needs [judge].model or REACT_JUDGE_MODEL"
|
|
47
|
+
),
|
|
48
|
+
warnings=list(self.load_warnings),
|
|
49
|
+
)
|
|
50
|
+
|
|
51
|
+
workspace = Path(self.cwd or os.getcwd()).resolve()
|
|
52
|
+
if not workspace.is_dir():
|
|
53
|
+
return AgentAttempt(
|
|
54
|
+
error=ValueError(f"judge workspace does not exist: {workspace}"),
|
|
55
|
+
warnings=list(self.load_warnings),
|
|
56
|
+
)
|
|
57
|
+
|
|
58
|
+
max_steps = _positive_int_env("REACT_JUDGE_MAX_STEPS", _DEFAULT_MAX_STEPS)
|
|
59
|
+
max_tool_chars = _positive_int_env(
|
|
60
|
+
"REACT_JUDGE_MAX_TOOL_CHARS", _DEFAULT_MAX_TOOL_CHARS
|
|
61
|
+
)
|
|
62
|
+
max_image_bytes = _positive_int_env(
|
|
63
|
+
"REACT_JUDGE_MAX_IMAGE_BYTES", _DEFAULT_MAX_IMAGE_BYTES
|
|
64
|
+
)
|
|
65
|
+
timeout = float(self.judge.timeout)
|
|
66
|
+
started = time.monotonic()
|
|
67
|
+
usage = _Usage(model)
|
|
68
|
+
trace: list[dict[str, Any]] = []
|
|
69
|
+
last_request: dict[str, Any] | None = None
|
|
70
|
+
last_response: dict[str, Any] | None = None
|
|
71
|
+
warnings = list(self.load_warnings)
|
|
72
|
+
|
|
73
|
+
messages: list[dict[str, Any]] = [
|
|
74
|
+
{
|
|
75
|
+
"role": "system",
|
|
76
|
+
"content": _system_prompt(workspace, max_steps),
|
|
77
|
+
},
|
|
78
|
+
{"role": "user", "content": prompt},
|
|
79
|
+
]
|
|
80
|
+
tools = [tool.function_schema() for tool in self.registry.values()]
|
|
81
|
+
tools.append(_submit_score_tool(schema))
|
|
82
|
+
|
|
83
|
+
try:
|
|
84
|
+
with tempfile.TemporaryDirectory(prefix="rewardkit-react-") as temp_dir:
|
|
85
|
+
async with asyncio.timeout(timeout):
|
|
86
|
+
for step in range(1, max_steps + 1):
|
|
87
|
+
remaining = max(0.1, timeout - (time.monotonic() - started))
|
|
88
|
+
completion_request = {
|
|
89
|
+
"model": model,
|
|
90
|
+
"messages": messages,
|
|
91
|
+
"tools": tools,
|
|
92
|
+
"timeout": remaining,
|
|
93
|
+
**_completion_overrides(),
|
|
94
|
+
}
|
|
95
|
+
last_request = _trace_request(completion_request)
|
|
96
|
+
last_response = None
|
|
97
|
+
response = await litellm.acompletion(
|
|
98
|
+
**completion_request,
|
|
99
|
+
)
|
|
100
|
+
usage.add_response(response)
|
|
101
|
+
message = response.choices[0].message
|
|
102
|
+
content = _field(message, "content")
|
|
103
|
+
tool_calls = list(_field(message, "tool_calls") or [])
|
|
104
|
+
assistant_message = _assistant_message(content, tool_calls)
|
|
105
|
+
last_response = copy.deepcopy(assistant_message)
|
|
106
|
+
messages.append(assistant_message)
|
|
107
|
+
trace.append(
|
|
108
|
+
{
|
|
109
|
+
"step": step,
|
|
110
|
+
"event": "assistant",
|
|
111
|
+
"content": content,
|
|
112
|
+
"tool_calls": assistant_message.get("tool_calls", []),
|
|
113
|
+
}
|
|
114
|
+
)
|
|
115
|
+
|
|
116
|
+
if not tool_calls:
|
|
117
|
+
trace.append(
|
|
118
|
+
{
|
|
119
|
+
"step": step,
|
|
120
|
+
"event": "missing_tool_call",
|
|
121
|
+
"content": _shorten(str(content or ""), 1000),
|
|
122
|
+
}
|
|
123
|
+
)
|
|
124
|
+
messages.append(
|
|
125
|
+
{
|
|
126
|
+
"role": "user",
|
|
127
|
+
"content": (
|
|
128
|
+
"Continue by calling one of the provided tools. "
|
|
129
|
+
"When the evidence is sufficient, call submit_score."
|
|
130
|
+
),
|
|
131
|
+
}
|
|
132
|
+
)
|
|
133
|
+
continue
|
|
134
|
+
|
|
135
|
+
image_observations: list[tuple[str, ToolResult]] = []
|
|
136
|
+
for tool_call in tool_calls:
|
|
137
|
+
call_id, name, arguments, argument_error = _parse_tool_call(
|
|
138
|
+
tool_call
|
|
139
|
+
)
|
|
140
|
+
usage.tool_requests[name] += 1
|
|
141
|
+
if argument_error:
|
|
142
|
+
result = ToolResult(error=argument_error)
|
|
143
|
+
elif name == "submit_score":
|
|
144
|
+
trace.append(
|
|
145
|
+
{
|
|
146
|
+
"step": step,
|
|
147
|
+
"tool": name,
|
|
148
|
+
"arguments": arguments,
|
|
149
|
+
}
|
|
150
|
+
)
|
|
151
|
+
return AgentAttempt(
|
|
152
|
+
output=json.dumps(arguments, ensure_ascii=False),
|
|
153
|
+
usage=usage.as_rewardkit(),
|
|
154
|
+
logs=_save_trace(
|
|
155
|
+
label=label,
|
|
156
|
+
model=model,
|
|
157
|
+
workspace=workspace,
|
|
158
|
+
outcome="completed",
|
|
159
|
+
request=last_request,
|
|
160
|
+
response=last_response,
|
|
161
|
+
trace=trace,
|
|
162
|
+
usage=usage.as_rewardkit(),
|
|
163
|
+
warnings=warnings,
|
|
164
|
+
),
|
|
165
|
+
warnings=warnings,
|
|
166
|
+
)
|
|
167
|
+
else:
|
|
168
|
+
result = await _execute_tool(
|
|
169
|
+
self.registry,
|
|
170
|
+
name,
|
|
171
|
+
arguments,
|
|
172
|
+
ToolContext(
|
|
173
|
+
workspace=workspace,
|
|
174
|
+
temp_dir=Path(temp_dir),
|
|
175
|
+
timeout=min(remaining, 120.0),
|
|
176
|
+
max_output_chars=max_tool_chars,
|
|
177
|
+
),
|
|
178
|
+
)
|
|
179
|
+
|
|
180
|
+
observation = result.observation(max_tool_chars)
|
|
181
|
+
messages.append(
|
|
182
|
+
{
|
|
183
|
+
"role": "tool",
|
|
184
|
+
"tool_call_id": call_id,
|
|
185
|
+
"name": name,
|
|
186
|
+
"content": observation,
|
|
187
|
+
}
|
|
188
|
+
)
|
|
189
|
+
trace.append(
|
|
190
|
+
{
|
|
191
|
+
"step": step,
|
|
192
|
+
"tool": name,
|
|
193
|
+
"arguments": arguments,
|
|
194
|
+
"observation": observation,
|
|
195
|
+
"images": [str(path) for path in result.images],
|
|
196
|
+
}
|
|
197
|
+
)
|
|
198
|
+
if result.images:
|
|
199
|
+
image_observations.append((name, result))
|
|
200
|
+
|
|
201
|
+
if image_observations:
|
|
202
|
+
blocks, image_warnings = _image_blocks(
|
|
203
|
+
image_observations, max_image_bytes
|
|
204
|
+
)
|
|
205
|
+
warnings.extend(image_warnings)
|
|
206
|
+
if blocks:
|
|
207
|
+
messages.append({"role": "user", "content": blocks})
|
|
208
|
+
|
|
209
|
+
except (TimeoutError, asyncio.TimeoutError, litellm.Timeout):
|
|
210
|
+
return AgentAttempt(
|
|
211
|
+
usage=usage.as_rewardkit(),
|
|
212
|
+
logs=_save_trace(
|
|
213
|
+
label=label,
|
|
214
|
+
model=model,
|
|
215
|
+
workspace=workspace,
|
|
216
|
+
outcome="timed_out",
|
|
217
|
+
request=last_request,
|
|
218
|
+
response=last_response,
|
|
219
|
+
trace=trace,
|
|
220
|
+
usage=usage.as_rewardkit(),
|
|
221
|
+
warnings=warnings,
|
|
222
|
+
),
|
|
223
|
+
warnings=warnings,
|
|
224
|
+
timed_out=True,
|
|
225
|
+
)
|
|
226
|
+
except Exception as exc:
|
|
227
|
+
trace.append(
|
|
228
|
+
{
|
|
229
|
+
"event": "error",
|
|
230
|
+
"error": f"{type(exc).__name__}: {exc}",
|
|
231
|
+
}
|
|
232
|
+
)
|
|
233
|
+
return AgentAttempt(
|
|
234
|
+
usage=usage.as_rewardkit(),
|
|
235
|
+
logs=_save_trace(
|
|
236
|
+
label=label,
|
|
237
|
+
model=model,
|
|
238
|
+
workspace=workspace,
|
|
239
|
+
outcome="error",
|
|
240
|
+
request=last_request,
|
|
241
|
+
response=last_response,
|
|
242
|
+
trace=trace,
|
|
243
|
+
usage=usage.as_rewardkit(),
|
|
244
|
+
warnings=warnings,
|
|
245
|
+
),
|
|
246
|
+
warnings=warnings,
|
|
247
|
+
error=exc,
|
|
248
|
+
)
|
|
249
|
+
|
|
250
|
+
trace.append(
|
|
251
|
+
{
|
|
252
|
+
"event": "error",
|
|
253
|
+
"error": f"maximum model turns reached: {max_steps}",
|
|
254
|
+
}
|
|
255
|
+
)
|
|
256
|
+
return AgentAttempt(
|
|
257
|
+
usage=usage.as_rewardkit(),
|
|
258
|
+
logs=_save_trace(
|
|
259
|
+
label=label,
|
|
260
|
+
model=model,
|
|
261
|
+
workspace=workspace,
|
|
262
|
+
outcome="max_steps",
|
|
263
|
+
request=last_request,
|
|
264
|
+
response=last_response,
|
|
265
|
+
trace=trace,
|
|
266
|
+
usage=usage.as_rewardkit(),
|
|
267
|
+
warnings=warnings,
|
|
268
|
+
),
|
|
269
|
+
warnings=warnings,
|
|
270
|
+
error=RuntimeError(
|
|
271
|
+
f"react judge reached {max_steps} model turns without submit_score"
|
|
272
|
+
),
|
|
273
|
+
)
|
|
274
|
+
|
|
275
|
+
|
|
276
|
+
async def _execute_tool(
|
|
277
|
+
registry: ToolRegistry,
|
|
278
|
+
name: str,
|
|
279
|
+
arguments: dict[str, Any],
|
|
280
|
+
context: ToolContext,
|
|
281
|
+
) -> ToolResult:
|
|
282
|
+
tool = registry.get(name)
|
|
283
|
+
if tool is None:
|
|
284
|
+
return ToolResult(error=f"unknown tool: {name}")
|
|
285
|
+
try:
|
|
286
|
+
return await tool.execute(arguments, context)
|
|
287
|
+
except Exception as exc:
|
|
288
|
+
return ToolResult(error=f"{type(exc).__name__}: {exc}")
|
|
289
|
+
|
|
290
|
+
|
|
291
|
+
def _system_prompt(workspace: Path, max_steps: int) -> str:
|
|
292
|
+
return f"""You are an evaluation agent judging artifacts in this workspace:
|
|
293
|
+
{workspace}
|
|
294
|
+
|
|
295
|
+
Use the provided tools to gather concrete evidence before scoring. Prefer targeted
|
|
296
|
+
inspection over reading every file. Do not modify the judged workspace. You have at
|
|
297
|
+
most {max_steps} model turns. When you have enough evidence, call submit_score exactly
|
|
298
|
+
once with the final scoring object. Do not merely print the score in assistant text;
|
|
299
|
+
the run ends only when you call submit_score. Include concise, evidence-based reasoning
|
|
300
|
+
in every criterion's reasoning field."""
|
|
301
|
+
|
|
302
|
+
|
|
303
|
+
def _submit_score_tool(schema: dict[str, Any]) -> dict[str, Any]:
|
|
304
|
+
return {
|
|
305
|
+
"type": "function",
|
|
306
|
+
"function": {
|
|
307
|
+
"name": "submit_score",
|
|
308
|
+
"description": (
|
|
309
|
+
"Submit the final RewardKit score after inspecting enough evidence. "
|
|
310
|
+
"Calling this tool ends the judge run."
|
|
311
|
+
),
|
|
312
|
+
"parameters": schema,
|
|
313
|
+
},
|
|
314
|
+
}
|
|
315
|
+
|
|
316
|
+
|
|
317
|
+
def _completion_overrides() -> dict[str, Any]:
|
|
318
|
+
result: dict[str, Any] = {}
|
|
319
|
+
base_url = os.environ.get("REACT_JUDGE_API_BASE")
|
|
320
|
+
api_key = os.environ.get("REACT_JUDGE_API_KEY")
|
|
321
|
+
if base_url:
|
|
322
|
+
result["base_url"] = base_url
|
|
323
|
+
if api_key:
|
|
324
|
+
result["api_key"] = api_key
|
|
325
|
+
return result
|
|
326
|
+
|
|
327
|
+
|
|
328
|
+
def _assistant_message(content: Any, tool_calls: list[Any]) -> dict[str, Any]:
|
|
329
|
+
result: dict[str, Any] = {"role": "assistant", "content": content}
|
|
330
|
+
if tool_calls:
|
|
331
|
+
result["tool_calls"] = []
|
|
332
|
+
for call in tool_calls:
|
|
333
|
+
function = _field(call, "function")
|
|
334
|
+
result["tool_calls"].append(
|
|
335
|
+
{
|
|
336
|
+
"id": str(_field(call, "id") or ""),
|
|
337
|
+
"type": "function",
|
|
338
|
+
"function": {
|
|
339
|
+
"name": str(_field(function, "name") or ""),
|
|
340
|
+
"arguments": _field(function, "arguments") or "{}",
|
|
341
|
+
},
|
|
342
|
+
}
|
|
343
|
+
)
|
|
344
|
+
return result
|
|
345
|
+
|
|
346
|
+
|
|
347
|
+
def _parse_tool_call(
|
|
348
|
+
call: Any,
|
|
349
|
+
) -> tuple[str, str, dict[str, Any], str | None]:
|
|
350
|
+
call_id = str(_field(call, "id") or f"call-{id(call)}")
|
|
351
|
+
function = _field(call, "function")
|
|
352
|
+
name = str(_field(function, "name") or "")
|
|
353
|
+
raw_arguments = _field(function, "arguments")
|
|
354
|
+
if isinstance(raw_arguments, dict):
|
|
355
|
+
arguments = raw_arguments
|
|
356
|
+
else:
|
|
357
|
+
try:
|
|
358
|
+
arguments = json.loads(raw_arguments or "{}")
|
|
359
|
+
except (TypeError, json.JSONDecodeError) as exc:
|
|
360
|
+
return call_id, name, {}, f"invalid JSON tool arguments: {exc}"
|
|
361
|
+
if not isinstance(arguments, dict):
|
|
362
|
+
return call_id, name, {}, "tool arguments must decode to a JSON object"
|
|
363
|
+
if not name:
|
|
364
|
+
return call_id, name, arguments, "tool call has no function name"
|
|
365
|
+
return call_id, name, arguments, None
|
|
366
|
+
|
|
367
|
+
|
|
368
|
+
def _field(value: Any, name: str, default: Any = None) -> Any:
|
|
369
|
+
if isinstance(value, dict):
|
|
370
|
+
return value.get(name, default)
|
|
371
|
+
return getattr(value, name, default)
|
|
372
|
+
|
|
373
|
+
|
|
374
|
+
def _image_blocks(
|
|
375
|
+
observations: list[tuple[str, ToolResult]], max_bytes: int
|
|
376
|
+
) -> tuple[list[dict[str, Any]], list[str]]:
|
|
377
|
+
blocks: list[dict[str, Any]] = []
|
|
378
|
+
warnings: list[str] = []
|
|
379
|
+
for tool_name, result in observations:
|
|
380
|
+
for path in result.images:
|
|
381
|
+
try:
|
|
382
|
+
data, mime_type = _read_image(path, max_bytes)
|
|
383
|
+
except Exception as exc:
|
|
384
|
+
warnings.append(f"could not attach image {path}: {exc}")
|
|
385
|
+
continue
|
|
386
|
+
blocks.append(
|
|
387
|
+
{
|
|
388
|
+
"type": "text",
|
|
389
|
+
"text": f"Visual observation from {tool_name}: {path.name}",
|
|
390
|
+
}
|
|
391
|
+
)
|
|
392
|
+
blocks.append(
|
|
393
|
+
{
|
|
394
|
+
"type": "image_url",
|
|
395
|
+
"image_url": {
|
|
396
|
+
"url": f"data:{mime_type};base64,{base64.b64encode(data).decode('ascii')}"
|
|
397
|
+
},
|
|
398
|
+
}
|
|
399
|
+
)
|
|
400
|
+
return blocks, warnings
|
|
401
|
+
|
|
402
|
+
|
|
403
|
+
def _read_image(path: Path, max_bytes: int) -> tuple[bytes, str]:
|
|
404
|
+
data = path.read_bytes()
|
|
405
|
+
mime_type = mimetypes.guess_type(path.name)[0] or "image/png"
|
|
406
|
+
if len(data) <= max_bytes:
|
|
407
|
+
return data, mime_type
|
|
408
|
+
|
|
409
|
+
try:
|
|
410
|
+
from PIL import Image
|
|
411
|
+
except ImportError as exc:
|
|
412
|
+
raise ValueError(
|
|
413
|
+
f"image is {len(data)} bytes (limit {max_bytes}) and Pillow is unavailable"
|
|
414
|
+
) from exc
|
|
415
|
+
|
|
416
|
+
with Image.open(path) as image:
|
|
417
|
+
image.thumbnail((1800, 1800))
|
|
418
|
+
if image.mode not in {"RGB", "L"}:
|
|
419
|
+
image = image.convert("RGB")
|
|
420
|
+
output = io.BytesIO()
|
|
421
|
+
image.save(output, format="JPEG", quality=85, optimize=True)
|
|
422
|
+
data = output.getvalue()
|
|
423
|
+
if len(data) > max_bytes:
|
|
424
|
+
raise ValueError(f"resized image is still larger than {max_bytes} bytes")
|
|
425
|
+
return data, "image/jpeg"
|
|
426
|
+
|
|
427
|
+
|
|
428
|
+
def _positive_int_env(name: str, default: int) -> int:
|
|
429
|
+
raw = os.environ.get(name)
|
|
430
|
+
if raw is None:
|
|
431
|
+
return default
|
|
432
|
+
value = int(raw)
|
|
433
|
+
if value <= 0:
|
|
434
|
+
raise ValueError(f"{name} must be positive")
|
|
435
|
+
return value
|
|
436
|
+
|
|
437
|
+
|
|
438
|
+
def _shorten(value: str, limit: int) -> str:
|
|
439
|
+
return value if len(value) <= limit else value[:limit] + "...[truncated]"
|
|
440
|
+
|
|
441
|
+
|
|
442
|
+
def _trace_request(request: dict[str, Any]) -> dict[str, Any]:
|
|
443
|
+
"""Copy the model request for the trace while excluding credentials."""
|
|
444
|
+
return {
|
|
445
|
+
key: copy.deepcopy(value)
|
|
446
|
+
for key, value in request.items()
|
|
447
|
+
if key.lower() not in {"api_key", "authorization"}
|
|
448
|
+
}
|
|
449
|
+
|
|
450
|
+
|
|
451
|
+
def _save_trace(
|
|
452
|
+
*,
|
|
453
|
+
label: str,
|
|
454
|
+
model: str,
|
|
455
|
+
workspace: Path,
|
|
456
|
+
outcome: str,
|
|
457
|
+
request: dict[str, Any] | None,
|
|
458
|
+
response: dict[str, Any] | None,
|
|
459
|
+
trace: list[dict[str, Any]],
|
|
460
|
+
usage: dict[str, Any],
|
|
461
|
+
warnings: list[str],
|
|
462
|
+
) -> list[str]:
|
|
463
|
+
configured = os.environ.get("REACT_JUDGE_LOG_DIR")
|
|
464
|
+
log_dir = (
|
|
465
|
+
Path(configured).expanduser()
|
|
466
|
+
if configured
|
|
467
|
+
else Path("/logs/verifier/judge_agents_trajs")
|
|
468
|
+
)
|
|
469
|
+
safe_label = "".join(
|
|
470
|
+
char if char.isalnum() or char in "-_" else "_" for char in label
|
|
471
|
+
)
|
|
472
|
+
path = log_dir / f"trajectory.react.{safe_label}.{uuid.uuid4().hex}.json"
|
|
473
|
+
payload = {
|
|
474
|
+
"version": 2,
|
|
475
|
+
"backend": "react",
|
|
476
|
+
"label": label,
|
|
477
|
+
"model": model,
|
|
478
|
+
"workspace": str(workspace),
|
|
479
|
+
"outcome": outcome,
|
|
480
|
+
"usage": usage,
|
|
481
|
+
"request": request,
|
|
482
|
+
"response": response,
|
|
483
|
+
"events": trace,
|
|
484
|
+
}
|
|
485
|
+
try:
|
|
486
|
+
log_dir.mkdir(parents=True, exist_ok=True)
|
|
487
|
+
path.write_text(
|
|
488
|
+
json.dumps(payload, ensure_ascii=False, indent=2, default=str) + "\n"
|
|
489
|
+
)
|
|
490
|
+
except Exception as exc:
|
|
491
|
+
warnings.append(f"could not save react judge trajectory to {path}: {exc}")
|
|
492
|
+
return []
|
|
493
|
+
return [str(path)]
|
|
494
|
+
|
|
495
|
+
|
|
496
|
+
class _Usage:
|
|
497
|
+
def __init__(self, model: str) -> None:
|
|
498
|
+
self.model = model
|
|
499
|
+
self.input_tokens = 0
|
|
500
|
+
self.output_tokens = 0
|
|
501
|
+
self.reasoning_output_tokens = 0
|
|
502
|
+
self.total_tokens = 0
|
|
503
|
+
self.tool_requests: Counter[str] = Counter()
|
|
504
|
+
|
|
505
|
+
def add_response(self, response: Any) -> None:
|
|
506
|
+
raw = _field(response, "usage")
|
|
507
|
+
if raw is None:
|
|
508
|
+
return
|
|
509
|
+
self.input_tokens += int(
|
|
510
|
+
_field(raw, "prompt_tokens", _field(raw, "input_tokens", 0)) or 0
|
|
511
|
+
)
|
|
512
|
+
self.output_tokens += int(
|
|
513
|
+
_field(raw, "completion_tokens", _field(raw, "output_tokens", 0)) or 0
|
|
514
|
+
)
|
|
515
|
+
self.total_tokens += int(_field(raw, "total_tokens", 0) or 0)
|
|
516
|
+
details = _field(raw, "completion_tokens_details") or {}
|
|
517
|
+
self.reasoning_output_tokens += int(_field(details, "reasoning_tokens", 0) or 0)
|
|
518
|
+
|
|
519
|
+
def as_rewardkit(self) -> dict[str, Any]:
|
|
520
|
+
result: dict[str, Any] = {}
|
|
521
|
+
if self.input_tokens:
|
|
522
|
+
result["input_tokens"] = self.input_tokens
|
|
523
|
+
if self.output_tokens:
|
|
524
|
+
result["output_tokens"] = self.output_tokens
|
|
525
|
+
if self.reasoning_output_tokens:
|
|
526
|
+
result["reasoning_output_tokens"] = self.reasoning_output_tokens
|
|
527
|
+
total = self.total_tokens or self.input_tokens + self.output_tokens
|
|
528
|
+
if total:
|
|
529
|
+
result["total_tokens"] = total
|
|
530
|
+
if self.tool_requests:
|
|
531
|
+
result["tool_requests"] = dict(self.tool_requests)
|
|
532
|
+
if result:
|
|
533
|
+
model_usage = {
|
|
534
|
+
key: value for key, value in result.items() if key != "models"
|
|
535
|
+
}
|
|
536
|
+
result["models"] = {self.model: model_usage}
|
|
537
|
+
return result
|
|
@@ -0,0 +1,152 @@
|
|
|
1
|
+
"""Tool registry and plugin discovery."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import importlib
|
|
6
|
+
import importlib.util
|
|
7
|
+
import os
|
|
8
|
+
from importlib.metadata import entry_points
|
|
9
|
+
from pathlib import Path
|
|
10
|
+
from types import ModuleType
|
|
11
|
+
from typing import Any, Iterable
|
|
12
|
+
|
|
13
|
+
from rewardkit_react.tool import Tool
|
|
14
|
+
|
|
15
|
+
ENTRY_POINT_GROUP = "rewardkit_react.tools"
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
class ToolRegistry:
|
|
19
|
+
def __init__(self) -> None:
|
|
20
|
+
self._tools: dict[str, Tool] = {}
|
|
21
|
+
|
|
22
|
+
def register(self, tool: Tool, *, replace: bool = False) -> None:
|
|
23
|
+
name = getattr(tool, "name", "")
|
|
24
|
+
if not isinstance(name, str) or not name:
|
|
25
|
+
raise TypeError("tool.name must be a non-empty string")
|
|
26
|
+
if name == "submit_score":
|
|
27
|
+
raise ValueError("submit_score is reserved by the agent")
|
|
28
|
+
if not callable(getattr(tool, "execute", None)):
|
|
29
|
+
raise TypeError(f"tool {name!r} has no execute() method")
|
|
30
|
+
if name in self._tools and not replace:
|
|
31
|
+
raise ValueError(f"tool already registered: {name}")
|
|
32
|
+
self._tools[name] = tool
|
|
33
|
+
|
|
34
|
+
def get(self, name: str) -> Tool | None:
|
|
35
|
+
return self._tools.get(name)
|
|
36
|
+
|
|
37
|
+
def values(self) -> tuple[Tool, ...]:
|
|
38
|
+
return tuple(self._tools[name] for name in sorted(self._tools))
|
|
39
|
+
|
|
40
|
+
def retain(self, names: set[str]) -> None:
|
|
41
|
+
unknown = names.difference(self._tools)
|
|
42
|
+
if unknown:
|
|
43
|
+
raise ValueError(f"unknown tools in REACT_JUDGE_TOOLS: {sorted(unknown)}")
|
|
44
|
+
self._tools = {
|
|
45
|
+
name: tool for name, tool in self._tools.items() if name in names
|
|
46
|
+
}
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
def _materialize(candidate: Any) -> Tool:
|
|
50
|
+
if isinstance(candidate, type):
|
|
51
|
+
candidate = candidate()
|
|
52
|
+
elif callable(candidate) and not hasattr(candidate, "execute"):
|
|
53
|
+
candidate = candidate()
|
|
54
|
+
if not hasattr(candidate, "name") or not callable(
|
|
55
|
+
getattr(candidate, "execute", None)
|
|
56
|
+
):
|
|
57
|
+
raise TypeError(
|
|
58
|
+
"plugin must be a Tool instance, Tool class, or zero-argument factory"
|
|
59
|
+
)
|
|
60
|
+
return candidate
|
|
61
|
+
|
|
62
|
+
|
|
63
|
+
def _register_collection(registry: ToolRegistry, values: Iterable[Any]) -> None:
|
|
64
|
+
for value in values:
|
|
65
|
+
registry.register(_materialize(value))
|
|
66
|
+
|
|
67
|
+
|
|
68
|
+
def _load_module_from_path(path: Path) -> ModuleType:
|
|
69
|
+
module_name = f"rewardkit_react_local_{abs(hash(path.resolve()))}"
|
|
70
|
+
spec = importlib.util.spec_from_file_location(module_name, path)
|
|
71
|
+
if spec is None or spec.loader is None:
|
|
72
|
+
raise ImportError(f"cannot load tool module: {path}")
|
|
73
|
+
module = importlib.util.module_from_spec(spec)
|
|
74
|
+
spec.loader.exec_module(module)
|
|
75
|
+
return module
|
|
76
|
+
|
|
77
|
+
|
|
78
|
+
def _register_module(registry: ToolRegistry, module: ModuleType) -> None:
|
|
79
|
+
hook = getattr(module, "register_tools", None)
|
|
80
|
+
if callable(hook):
|
|
81
|
+
hook(registry)
|
|
82
|
+
return
|
|
83
|
+
factory = getattr(module, "get_tools", None)
|
|
84
|
+
if callable(factory):
|
|
85
|
+
_register_collection(registry, factory())
|
|
86
|
+
return
|
|
87
|
+
tools = getattr(module, "TOOLS", None)
|
|
88
|
+
if tools is not None:
|
|
89
|
+
_register_collection(registry, tools)
|
|
90
|
+
return
|
|
91
|
+
raise ValueError(
|
|
92
|
+
f"{module.__name__} must expose register_tools(registry), get_tools(), or TOOLS"
|
|
93
|
+
)
|
|
94
|
+
|
|
95
|
+
|
|
96
|
+
def load_registry() -> tuple[ToolRegistry, list[str]]:
|
|
97
|
+
"""Load built-ins, installed entry points, and task-local modules."""
|
|
98
|
+
from rewardkit_react.tools import (
|
|
99
|
+
InspectFileTool,
|
|
100
|
+
ListFilesTool,
|
|
101
|
+
RenderDocumentTool,
|
|
102
|
+
SearchFilesTool,
|
|
103
|
+
ShellTool,
|
|
104
|
+
ViewTool,
|
|
105
|
+
)
|
|
106
|
+
|
|
107
|
+
registry = ToolRegistry()
|
|
108
|
+
for tool in (
|
|
109
|
+
ShellTool(),
|
|
110
|
+
ViewTool(),
|
|
111
|
+
ListFilesTool(),
|
|
112
|
+
SearchFilesTool(),
|
|
113
|
+
InspectFileTool(),
|
|
114
|
+
RenderDocumentTool(),
|
|
115
|
+
):
|
|
116
|
+
registry.register(tool)
|
|
117
|
+
|
|
118
|
+
warnings: list[str] = []
|
|
119
|
+
for entry_point in entry_points(group=ENTRY_POINT_GROUP):
|
|
120
|
+
try:
|
|
121
|
+
registry.register(_materialize(entry_point.load()))
|
|
122
|
+
except Exception as exc:
|
|
123
|
+
warnings.append(
|
|
124
|
+
f"failed to load tool entry point {entry_point.name}: {exc}"
|
|
125
|
+
)
|
|
126
|
+
|
|
127
|
+
raw_modules = os.environ.get("REACT_JUDGE_TOOL_MODULES", "")
|
|
128
|
+
module_names = [
|
|
129
|
+
item.strip() for item in raw_modules.replace(",", os.pathsep).split(os.pathsep)
|
|
130
|
+
]
|
|
131
|
+
for module_name in filter(None, module_names):
|
|
132
|
+
try:
|
|
133
|
+
path = Path(module_name)
|
|
134
|
+
module = (
|
|
135
|
+
_load_module_from_path(path)
|
|
136
|
+
if path.suffix == ".py" or path.exists()
|
|
137
|
+
else importlib.import_module(module_name)
|
|
138
|
+
)
|
|
139
|
+
_register_module(registry, module)
|
|
140
|
+
except Exception as exc:
|
|
141
|
+
warnings.append(f"failed to load tool module {module_name}: {exc}")
|
|
142
|
+
|
|
143
|
+
enabled = os.environ.get("REACT_JUDGE_TOOLS")
|
|
144
|
+
if enabled:
|
|
145
|
+
registry.retain({name.strip() for name in enabled.split(",") if name.strip()})
|
|
146
|
+
elif os.environ.get("REACT_JUDGE_ENABLE_SHELL", "").lower() not in {
|
|
147
|
+
"1",
|
|
148
|
+
"true",
|
|
149
|
+
"yes",
|
|
150
|
+
}:
|
|
151
|
+
registry.retain({tool.name for tool in registry.values()} - {"shell"})
|
|
152
|
+
return registry, warnings
|