Use single-tool required choice for reliable forced search arguments

This commit is contained in:
pewdiepie-archdaemon
2026-09-17 21:31:09 +00:00
parent a23d709056
commit a80a090d37
4 changed files with 59 additions and 8 deletions
+18 -2
View File
@@ -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 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 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 = []; 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), '<canonical-system-expression>', 'eval')))
`], {input:fs.readFileSync('src/clean_agent_preview.py','utf8'),encoding:'utf8'}).trim();
}
for (const prompt of prompts) { for (const prompt of prompts) {
for (const choice of ['auto', 'required', {type:'function', function:{name:'web_search'}}]) { for (const choice of ['auto', 'required', {type:'function', function:{name:'web_search'}}]) {
const started = performance.now(); const started = performance.now();
const response = await fetch(endpoint, {method:'POST', headers:{'Content-Type':'application/json'}, signal:AbortSignal.timeout(90000), 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 data = await response.json();
const result = {prompt,choice,status:response.status,seconds:(performance.now()-started)/1000,message:data.choices?.[0]?.message,error:data.error}; 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)); results.push(result); console.log(JSON.stringify(result));
} }
} }
const target = `reports/search-tool-choice-probe-${Date.now()}.json`; 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); console.log(target);
+20
View File
@@ -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): def bounded_search_observation(output, budget=8000):
"""Share the observation budget across fetched sources, not prefix order. """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 force_private_browser_next_round = False
pending, content = {}, '' pending, content = {}, ''
request = search_tool_choice_request(request)
async with preview_model_response(client, endpoint_url, headers, request, context_recovery) as response: async with preview_model_response(client, endpoint_url, headers, request, context_recovery) as response:
response.raise_for_status() response.raise_for_status()
async for line in response.aiter_lines(): async for line in response.aiter_lines():
+4 -6
View File
@@ -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] 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 len(requests) == (4 if embedded_article or empty_second_search else 5)
assert requests[2]['tool_choice'] == { assert requests[2]['tool_choice'] == 'required'
'type': 'function', 'function': {'name': 'web_search'}, assert [s['function']['name'] for s in requests[2]['tools']] == ['web_search']
}
if not embedded_article and not empty_second_search: if not embedded_article and not empty_second_search:
assert requests[3]['tool_choice'] == { assert requests[3]['tool_choice'] == {
'type': 'function', 'function': {'name': 'web_fetch'}, '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 executions == ['private_browser', 'web_search']
assert requests[1]['tool_choice'] == { assert requests[1]['tool_choice'] == 'required'
'type': 'function', 'function': {'name': 'web_search'}, 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] events = [json.loads(chunk[6:]) for chunk in raw if '[DONE]' not in chunk]
assert any( assert any(
event.get('type') == 'tool_loop_recovery' event.get('type') == 'tool_loop_recovery'
+17
View File
@@ -4,6 +4,23 @@ import json
import pytest 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', [ @pytest.mark.parametrize('prompt', [
'Explain the settings and link the instructions, not just the homepage.', 'Explain the settings and link the instructions, not just the homepage.',
'Can you link to the original studies?', 'Can you link to the original studies?',