diff --git a/scripts/probe_search_tool_choice.mjs b/scripts/probe_search_tool_choice.mjs index 2e1084760..83e133551 100644 --- a/scripts/probe_search_tool_choice.mjs +++ b/scripts/probe_search_tool_choice.mjs @@ -8,16 +8,32 @@ const model = process.env.MODEL || 'model-f'; const endpoint = process.env.ENDPOINT_URL || (() => { throw new Error("ENDPOINT_URL is required"); })(); const prompts = ['Catch me up on the biggest AI developments this week. Explain why they matter and link your sources.', 'serch latest ai news pls']; const results = []; +let system = 'You are Odysseus. Use web_search to find current information relevant to the user request.'; +if (process.env.CANONICAL_SYSTEM === '1') { + system = execFileSync((process.env.PYTHON || "python3"), ['-c', ` +import ast, sys +from datetime import datetime, timezone +from src.clean_agent_preview import native_input_files_clause +tree = ast.parse(sys.stdin.read()) +fn = next(n for n in tree.body if isinstance(n, ast.AsyncFunctionDef) and n.name == 'stream_preview') +assignment = next(n for n in fn.body if isinstance(n, ast.Assign) and any(isinstance(t, ast.Name) and t.id == 'system' for t in n.targets)) +runtime_scope_clause = 'This is a tool preview connected to the authenticated user’s real data. ' +native_workspace_enabled = False +client_runtime_context = None +shell_clause = 'Shell commands are disabled. ' +print(eval(compile(ast.Expression(assignment.value), '', 'eval'))) +`], {input:fs.readFileSync('src/clean_agent_preview.py','utf8'),encoding:'utf8'}).trim(); +} for (const prompt of prompts) { for (const choice of ['auto', 'required', {type:'function', function:{name:'web_search'}}]) { const started = performance.now(); const response = await fetch(endpoint, {method:'POST', headers:{'Content-Type':'application/json'}, signal:AbortSignal.timeout(90000), - body:JSON.stringify({model, messages:[{role:'system',content:'You are Odysseus. Use web_search to find current information relevant to the user request.'},{role:'user',content:prompt}], tools, tool_choice:choice, temperature:0, max_tokens:256, stream:false, chat_template_kwargs:{enable_thinking:false}})}); + body:JSON.stringify({model, messages:[{role:'system',content:system},{role:'user',content:prompt}], tools, tool_choice:choice, temperature:0, max_tokens:256, stream:false, chat_template_kwargs:{enable_thinking:false}})}); const data = await response.json(); const result = {prompt,choice,status:response.status,seconds:(performance.now()-started)/1000,message:data.choices?.[0]?.message,error:data.error}; results.push(result); console.log(JSON.stringify(result)); } } const target = `reports/search-tool-choice-probe-${Date.now()}.json`; -fs.writeFileSync(target, JSON.stringify({model,tools,results},null,2)+'\n'); +fs.writeFileSync(target, JSON.stringify({model,system,tools,results},null,2)+'\n'); console.log(target); diff --git a/src/clean_agent_preview.py b/src/clean_agent_preview.py index 840b674b0..ccc003345 100644 --- a/src/clean_agent_preview.py +++ b/src/clean_agent_preview.py @@ -162,6 +162,25 @@ SAFE_ACTIONS = { } +def search_tool_choice_request(request): + """Enforce a search via one offered tool, not named-tool argument decoding. + + The served model emits missing query fields under named search choice. + Required choice over the same single schema preserves the policy intent. + Other tools and auto/none requests retain their existing dispatch. + """ + choice = request.get('tool_choice') + if not isinstance(choice, dict) or choice.get('type') != 'function': + return request + name = (choice.get('function') or {}).get('name') + if name != 'web_search': + return request + selected = [s for s in request.get('tools', []) if s.get('function', {}).get('name') == name] + if len(selected) != 1: + return request + return {**request, 'tools': selected, 'tool_choice': 'required'} + + def bounded_search_observation(output, budget=8000): """Share the observation budget across fetched sources, not prefix order. @@ -4190,6 +4209,7 @@ async def stream_preview(*, endpoint_url, model, messages, headers, turn_contrac } force_private_browser_next_round = False pending, content = {}, '' + request = search_tool_choice_request(request) async with preview_model_response(client, endpoint_url, headers, request, context_recovery) as response: response.raise_for_status() async for line in response.aiter_lines(): diff --git a/tests/test_clean_agent_preview.py b/tests/test_clean_agent_preview.py index 6639d5fe4..9cb2d912b 100644 --- a/tests/test_clean_agent_preview.py +++ b/tests/test_clean_agent_preview.py @@ -1283,9 +1283,8 @@ async def test_stream_retries_an_obviously_truncated_broad_web_answer(monkeypatc events = [json.loads(chunk[6:]) for chunk in raw if '[DONE]' not in chunk] assert len(requests) == (4 if embedded_article or empty_second_search else 5) - assert requests[2]['tool_choice'] == { - 'type': 'function', 'function': {'name': 'web_search'}, - } + assert requests[2]['tool_choice'] == 'required' + assert [s['function']['name'] for s in requests[2]['tools']] == ['web_search'] if not embedded_article and not empty_second_search: assert requests[3]['tool_choice'] == { 'type': 'function', 'function': {'name': 'web_fetch'}, @@ -4351,9 +4350,8 @@ async def test_blocked_search_engine_browser_forces_native_web_search(monkeypatc )] assert executions == ['private_browser', 'web_search'] - assert requests[1]['tool_choice'] == { - 'type': 'function', 'function': {'name': 'web_search'}, - } + assert requests[1]['tool_choice'] == 'required' + assert [s['function']['name'] for s in requests[1]['tools']] == ['web_search'] events = [json.loads(chunk[6:]) for chunk in raw if '[DONE]' not in chunk] assert any( event.get('type') == 'tool_loop_recovery' diff --git a/tests/test_search_observation_budget.py b/tests/test_search_observation_budget.py index ba7a222ea..cda66631d 100644 --- a/tests/test_search_observation_budget.py +++ b/tests/test_search_observation_budget.py @@ -4,6 +4,23 @@ import json import pytest +def test_forced_search_dispatch_preserves_schema_without_mutating_request(): + from src.clean_agent_preview import search_tool_choice_request + search = {'type': 'function', 'function': {'name': 'web_search', 'parameters': {'required': ['query']}}} + other = {'type': 'function', 'function': {'name': 'web_fetch'}} + request = {'tools': [search, other], 'tool_choice': {'type': 'function', 'function': {'name': 'web_search'}}, 'messages': []} + converted = search_tool_choice_request(request) + assert converted['tools'] == [search] + assert converted['tools'][0] is search + assert converted['tool_choice'] == 'required' + assert len(request['tools']) == 2 + for choice in ['auto', 'none', 'required', {'type': 'function', 'function': {'name': 'web_fetch'}}]: + other_request = {**request, 'tool_choice': choice} + assert search_tool_choice_request(other_request) is other_request + missing = {**request, 'tools': [other]} + assert search_tool_choice_request(missing) is missing + + @pytest.mark.parametrize('prompt', [ 'Explain the settings and link the instructions, not just the homepage.', 'Can you link to the original studies?',