55
66import argparse
77import asyncio
8+ import datetime
89import json
9- import os
1010import pathlib
1111import re
1212import sys
@@ -47,7 +47,48 @@ def load_json(value: str | None, *, field: str, failures: list[dict]) -> object:
4747 return {}
4848
4949
50- async def audit (repo : pathlib .Path , run_id : str | None ) -> dict :
50+ def parse_created_after (value : str | None ) -> datetime .datetime | None :
51+ if not value :
52+ return None
53+ parsed = datetime .datetime .fromisoformat (value .replace ("Z" , "+00:00" ))
54+ if parsed .tzinfo is not None :
55+ parsed = parsed .astimezone (datetime .timezone .utc ).replace (tzinfo = None )
56+ return parsed
57+
58+
59+ def event_matches_tool_call (data_json : str | None , tool_name : str , parameters : dict | None ) -> bool :
60+ try :
61+ data = json .loads (data_json or "{}" )
62+ except (TypeError , ValueError ):
63+ return False
64+ if not isinstance (data , dict ) or data .get ("tool_name" ) != tool_name :
65+ return False
66+ return parameters is None or data .get ("parameters" ) == parameters
67+
68+
69+ def collect_result_texts (value : object ) -> list [str ]:
70+ texts : list [str ] = []
71+ if isinstance (value , dict ):
72+ for key , item in value .items ():
73+ if key == "text" and isinstance (item , str ):
74+ texts .append (item )
75+ else :
76+ texts .extend (collect_result_texts (item ))
77+ elif isinstance (value , list ):
78+ for item in value :
79+ texts .extend (collect_result_texts (item ))
80+ return texts
81+
82+
83+ async def audit (
84+ repo : pathlib .Path ,
85+ run_id : str | None ,
86+ * ,
87+ created_after : datetime .datetime | None = None ,
88+ expected_tool_name : str | None = None ,
89+ expected_parameters : dict | None = None ,
90+ expected_result_text : str | None = None ,
91+ ) -> dict :
5192 engine = create_async_engine (database_url (repo ))
5293 failures : list [dict ] = []
5394 warnings : list [dict ] = []
@@ -58,12 +99,41 @@ async def audit(repo: pathlib.Path, run_id: str | None) -> dict:
5899 sqlalchemy .text ("SELECT * FROM agent_run WHERE run_id = :run_id" ),
59100 {"run_id" : run_id },
60101 )).mappings ().first ()
102+ elif expected_tool_name :
103+ query = "SELECT * FROM agent_run"
104+ params = {}
105+ if created_after is not None :
106+ query += " WHERE created_at >= :created_after"
107+ params ["created_after" ] = created_after
108+ query += " ORDER BY id DESC LIMIT 100"
109+ candidates = (await connection .execute (sqlalchemy .text (query ), params )).mappings ().all ()
110+ run_row = None
111+ for candidate in candidates :
112+ started_rows = (await connection .execute (
113+ sqlalchemy .text (
114+ "SELECT data_json FROM agent_run_event "
115+ "WHERE run_id = :run_id AND type = 'tool.call.started' ORDER BY sequence"
116+ ),
117+ {"run_id" : str (candidate ["run_id" ])},
118+ )).mappings ().all ()
119+ if any (
120+ event_matches_tool_call (row .get ("data_json" ), expected_tool_name , expected_parameters )
121+ for row in started_rows
122+ ):
123+ run_row = candidate
124+ break
61125 else :
62126 run_row = (await connection .execute (
63127 sqlalchemy .text ("SELECT * FROM agent_run ORDER BY id DESC LIMIT 1" )
64128 )).mappings ().first ()
65129 if run_row is None :
66- return {"status" : "env_issue" , "reason" : "No matching AgentRunner run exists." , "failures" : [], "warnings" : []}
130+ status = "fail" if expected_tool_name else "env_issue"
131+ return {
132+ "status" : status ,
133+ "reason" : "No AgentRunner run contains the expected tool call." if expected_tool_name else "No matching AgentRunner run exists." ,
134+ "failures" : [{"kind" : "expected_tool_call_missing" }] if expected_tool_name else [],
135+ "warnings" : [],
136+ }
67137
68138 selected_run_id = str (run_row ["run_id" ])
69139 event_rows = (await connection .execute (
@@ -173,6 +243,36 @@ def error_surface(value: object) -> list[str]:
173243 if not tools :
174244 warnings .append ({"kind" : "no_authorized_tools" , "reason" : "The run authorization snapshot exposes no tools." })
175245
246+ expected_call_summary = None
247+ if expected_tool_name :
248+ matching_starts = [
249+ item
250+ for items in starts .values ()
251+ for item in items
252+ if item ["tool_name" ] == expected_tool_name
253+ and (expected_parameters is None or item ["data" ].get ("parameters" ) == expected_parameters )
254+ ]
255+ if len (matching_starts ) != 1 :
256+ failures .append ({"kind" : "expected_tool_call_count" , "actual" : len (matching_starts ), "expected" : 1 })
257+ matching_completions = []
258+ for started in matching_starts :
259+ call_id = str (started ["data" ].get ("tool_call_id" , "" ))
260+ matching_completions .extend (completions .get (call_id , []))
261+ result_text_match = expected_result_text is None or any (
262+ expected_result_text in collect_result_texts (completed ["data" ].get ("result" ))
263+ for completed in matching_completions
264+ )
265+ if expected_result_text is not None and not result_text_match :
266+ failures .append ({"kind" : "expected_tool_result_text_missing" })
267+ expected_call_summary = {
268+ "tool_name" : expected_tool_name ,
269+ "parameters_match_required" : expected_parameters is not None ,
270+ "matched_started_count" : len (matching_starts ),
271+ "matched_completed_count" : len (matching_completions ),
272+ "result_text_match_required" : expected_result_text is not None ,
273+ "result_text_match" : result_text_match ,
274+ }
275+
176276 metrics = {
177277 "event_count" : len (event_rows ),
178278 "tool_call_started" : sum (len (items ) for items in starts .values ()),
@@ -193,6 +293,7 @@ def error_surface(value: object) -> list[str]:
193293 "finished_at" : str (run_row ["finished_at" ]),
194294 },
195295 "metrics" : metrics ,
296+ "expected_tool_call" : expected_call_summary ,
196297 "failures" : failures ,
197298 "warnings" : warnings ,
198299 }
@@ -202,10 +303,28 @@ def main() -> int:
202303 parser = argparse .ArgumentParser ()
203304 parser .add_argument ("--repo" , required = True )
204305 parser .add_argument ("--run-id" )
306+ parser .add_argument ("--created-after" )
307+ parser .add_argument ("--expected-tool-name" )
308+ parser .add_argument ("--expected-parameters-json" )
309+ parser .add_argument ("--expected-result-text" )
205310 parser .add_argument ("--output" , required = True )
206311 args = parser .parse_args ()
207312 try :
208- report = asyncio .run (audit (pathlib .Path (args .repo ).resolve (), args .run_id ))
313+ expected_parameters = None
314+ if args .expected_parameters_json :
315+ expected_parameters = json .loads (args .expected_parameters_json )
316+ if not isinstance (expected_parameters , dict ):
317+ raise ValueError ("--expected-parameters-json must decode to an object" )
318+ if (expected_parameters is not None or args .expected_result_text ) and not args .expected_tool_name :
319+ raise ValueError ("--expected-tool-name is required with expected parameters or result text" )
320+ report = asyncio .run (audit (
321+ pathlib .Path (args .repo ).resolve (),
322+ args .run_id ,
323+ created_after = parse_created_after (args .created_after ),
324+ expected_tool_name = args .expected_tool_name ,
325+ expected_parameters = expected_parameters ,
326+ expected_result_text = args .expected_result_text ,
327+ ))
209328 except Exception as exc : # noqa: BLE001 - probe must classify environment failures
210329 report = {"status" : "env_issue" , "reason" : str (exc ), "failures" : [], "warnings" : []}
211330 pathlib .Path (args .output ).write_text (json .dumps (report , indent = 2 ) + "\n " , encoding = "utf-8" )
0 commit comments