From d8d2782e060576be5ec401e0c81e97faf67a8536 Mon Sep 17 00:00:00 2001 From: Lengshuang <90967079+Lesereingrape@users.noreply.github.com> Date: Sat, 26 Sep 2026 15:31:09 +0800 Subject: [PATCH] fix(infer): raise the intended ValueError when tool_choice names an unknown tool --- swift/infer_engine/protocol.py | 2 +- tests/general/test_protocol_tool_choice.py | 43 ++++++++++++++++++++++ 2 files changed, 44 insertions(+), 1 deletion(-) create mode 100644 tests/general/test_protocol_tool_choice.py diff --git a/swift/infer_engine/protocol.py b/swift/infer_engine/protocol.py index 68906a84ef..88bf48f518 100644 --- a/swift/infer_engine/protocol.py +++ b/swift/infer_engine/protocol.py @@ -247,7 +247,7 @@ def __post_init__(self): self.tools = None elif isinstance(self.tool_choice, dict): name = self.tool_choice['function']['name'] - tool = next(tool for tool in self.tools if tool['function']['name'] == name) + tool = next((tool for tool in self.tools if tool['function']['name'] == name), None) if tool is None: raise ValueError(f"Tool choice '{name}' not found in tools.") self.tools = [tool] diff --git a/tests/general/test_protocol_tool_choice.py b/tests/general/test_protocol_tool_choice.py new file mode 100644 index 0000000000..2e668f8e91 --- /dev/null +++ b/tests/general/test_protocol_tool_choice.py @@ -0,0 +1,43 @@ +import unittest + +from swift.infer_engine.protocol import ChatCompletionRequest + + +def _request(tool_choice): + tools = [ + { + 'type': 'function', + 'function': { + 'name': 'get_weather' + } + }, + { + 'type': 'function', + 'function': { + 'name': 'get_time' + } + }, + ] + return ChatCompletionRequest( + model='test', messages=[{ + 'role': 'user', + 'content': 'hi' + }], tools=tools, tool_choice=tool_choice) + + +class TestChatCompletionRequestToolChoice(unittest.TestCase): + + def test_selected_tool_keeps_only_that_tool(self): + request = _request({'type': 'function', 'function': {'name': 'get_time'}}) + self.assertEqual([tool['function']['name'] for tool in request.tools], ['get_time']) + + def test_unknown_tool_raises_value_error(self): + with self.assertRaisesRegex(ValueError, "Tool choice 'no_such_tool' not found in tools."): + _request({'type': 'function', 'function': {'name': 'no_such_tool'}}) + + def test_string_tool_choice_keeps_every_tool(self): + self.assertEqual(len(_request('auto').tools), 2) + + +if __name__ == '__main__': + unittest.main()