diff --git a/ollama/_types.py b/ollama/_types.py index b9bb53cb..de602173 100644 --- a/ollama/_types.py +++ b/ollama/_types.py @@ -390,8 +390,9 @@ def serialize_model(self, nxt): output = nxt(self) if output.get('tools'): for tool in output['tools']: - if 'function' in tool and 'parameters' in tool['function'] and 'defs' in tool['function']['parameters']: - tool['function']['parameters']['$defs'] = tool['function']['parameters'].pop('defs') + parameters = (tool.get('function') or {}).get('parameters') + if parameters and 'defs' in parameters: + parameters['$defs'] = parameters.pop('defs') return output messages: Optional[Sequence[Union[Mapping[str, Any], Message]]] = None diff --git a/tests/test_type_serialization.py b/tests/test_type_serialization.py index 02a69e95..405cbefa 100644 --- a/tests/test_type_serialization.py +++ b/tests/test_type_serialization.py @@ -4,7 +4,7 @@ import pytest -from ollama._types import CreateRequest, Image +from ollama._types import ChatRequest, CreateRequest, Image, Tool def test_image_serialization_bytes(): @@ -105,3 +105,40 @@ def test_create_request_serialization_license_list(): request = CreateRequest(model='test-model', license=['MIT', 'Apache-2.0']) serialized = request.model_dump() assert serialized['license'] == ['MIT', 'Apache-2.0'] + + +def test_chat_request_serialization_tool_without_parameters(): + request = ChatRequest(model='test-model', tools=[Tool(function=Tool.Function(name='func'))]) + serialized = request.model_dump() + assert serialized['tools'][0]['function']['name'] == 'func' + + serialized_json = request.model_dump_json() + assert 'func' in serialized_json + + +def test_chat_request_serialization_tool_without_function(): + request = ChatRequest(model='test-model', tools=[Tool()]) + serialized = request.model_dump() + assert serialized['tools'][0]['type'] == 'function' + + request.model_dump_json() + + +def test_chat_request_serialization_tool_defs_alias(): + request = ChatRequest( + model='test-model', + tools=[ + Tool( + function=Tool.Function( + name='func', + parameters=Tool.Function.Parameters( + **{'$defs': {'Nested': {'type': 'object'}}}, + properties={'arg': Tool.Function.Parameters.Property(type='string')}, + ), + ) + ) + ], + ) + serialized = request.model_dump() + assert serialized['tools'][0]['function']['parameters']['$defs'] == {'Nested': {'type': 'object'}} + assert 'defs' not in serialized['tools'][0]['function']['parameters']