diff --git a/ollama/_client.py b/ollama/_client.py index 8dfce824..37a697b7 100644 --- a/ollama/_client.py +++ b/ollama/_client.py @@ -209,7 +209,7 @@ def generate( template: str = '', context: Optional[Sequence[int]] = None, stream: Literal[False] = False, - think: Optional[Union[bool, Literal['low', 'medium', 'high']]] = None, + think: Optional[Union[bool, str]] = None, logprobs: Optional[bool] = None, top_logprobs: Optional[int] = None, raw: bool = False, @@ -233,7 +233,7 @@ def generate( template: str = '', context: Optional[Sequence[int]] = None, stream: Literal[True] = True, - think: Optional[Union[bool, Literal['low', 'medium', 'high']]] = None, + think: Optional[Union[bool, str]] = None, logprobs: Optional[bool] = None, top_logprobs: Optional[int] = None, raw: bool = False, @@ -256,7 +256,7 @@ def generate( template: Optional[str] = None, context: Optional[Sequence[int]] = None, stream: bool = False, - think: Optional[Union[bool, Literal['low', 'medium', 'high']]] = None, + think: Optional[Union[bool, str]] = None, logprobs: Optional[bool] = None, top_logprobs: Optional[int] = None, raw: Optional[bool] = None, @@ -313,7 +313,7 @@ def chat( *, tools: Optional[Sequence[Union[Mapping[str, Any], Tool, Callable]]] = None, stream: Literal[False] = False, - think: Optional[Union[bool, Literal['low', 'medium', 'high']]] = None, + think: Optional[Union[bool, str]] = None, logprobs: Optional[bool] = None, top_logprobs: Optional[int] = None, format: Optional[Union[Literal['', 'json'], JsonSchemaValue]] = None, @@ -329,7 +329,7 @@ def chat( *, tools: Optional[Sequence[Union[Mapping[str, Any], Tool, Callable]]] = None, stream: Literal[True] = True, - think: Optional[Union[bool, Literal['low', 'medium', 'high']]] = None, + think: Optional[Union[bool, str]] = None, logprobs: Optional[bool] = None, top_logprobs: Optional[int] = None, format: Optional[Union[Literal['', 'json'], JsonSchemaValue]] = None, @@ -344,7 +344,7 @@ def chat( *, tools: Optional[Sequence[Union[Mapping[str, Any], Tool, Callable]]] = None, stream: bool = False, - think: Optional[Union[bool, Literal['low', 'medium', 'high']]] = None, + think: Optional[Union[bool, str]] = None, logprobs: Optional[bool] = None, top_logprobs: Optional[int] = None, format: Optional[Union[Literal['', 'json'], JsonSchemaValue]] = None, @@ -842,7 +842,7 @@ async def generate( template: str = '', context: Optional[Sequence[int]] = None, stream: Literal[False] = False, - think: Optional[Union[bool, Literal['low', 'medium', 'high']]] = None, + think: Optional[Union[bool, str]] = None, logprobs: Optional[bool] = None, top_logprobs: Optional[int] = None, raw: bool = False, @@ -866,7 +866,7 @@ async def generate( template: str = '', context: Optional[Sequence[int]] = None, stream: Literal[True] = True, - think: Optional[Union[bool, Literal['low', 'medium', 'high']]] = None, + think: Optional[Union[bool, str]] = None, logprobs: Optional[bool] = None, top_logprobs: Optional[int] = None, raw: bool = False, @@ -889,7 +889,7 @@ async def generate( template: Optional[str] = None, context: Optional[Sequence[int]] = None, stream: bool = False, - think: Optional[Union[bool, Literal['low', 'medium', 'high']]] = None, + think: Optional[Union[bool, str]] = None, logprobs: Optional[bool] = None, top_logprobs: Optional[int] = None, raw: Optional[bool] = None, @@ -945,7 +945,7 @@ async def chat( *, tools: Optional[Sequence[Union[Mapping[str, Any], Tool, Callable]]] = None, stream: Literal[False] = False, - think: Optional[Union[bool, Literal['low', 'medium', 'high']]] = None, + think: Optional[Union[bool, str]] = None, logprobs: Optional[bool] = None, top_logprobs: Optional[int] = None, format: Optional[Union[Literal['', 'json'], JsonSchemaValue]] = None, @@ -961,7 +961,7 @@ async def chat( *, tools: Optional[Sequence[Union[Mapping[str, Any], Tool, Callable]]] = None, stream: Literal[True] = True, - think: Optional[Union[bool, Literal['low', 'medium', 'high']]] = None, + think: Optional[Union[bool, str]] = None, logprobs: Optional[bool] = None, top_logprobs: Optional[int] = None, format: Optional[Union[Literal['', 'json'], JsonSchemaValue]] = None, @@ -976,7 +976,7 @@ async def chat( *, tools: Optional[Sequence[Union[Mapping[str, Any], Tool, Callable]]] = None, stream: bool = False, - think: Optional[Union[bool, Literal['low', 'medium', 'high']]] = None, + think: Optional[Union[bool, str]] = None, logprobs: Optional[bool] = None, top_logprobs: Optional[int] = None, format: Optional[Union[Literal['', 'json'], JsonSchemaValue]] = None, diff --git a/ollama/_types.py b/ollama/_types.py index b9bb53cb..5931402c 100644 --- a/ollama/_types.py +++ b/ollama/_types.py @@ -207,7 +207,7 @@ class GenerateRequest(BaseGenerateRequest): images: Optional[Sequence[Image]] = None 'Image data for multimodal models.' - think: Optional[Union[bool, Literal['low', 'medium', 'high']]] = None + think: Optional[Union[bool, str]] = None 'Enable thinking mode (for thinking models).' logprobs: Optional[bool] = None @@ -400,7 +400,7 @@ def serialize_model(self, nxt): tools: Optional[Sequence[Tool]] = None 'Tools to use for the chat.' - think: Optional[Union[bool, Literal['low', 'medium', 'high']]] = None + think: Optional[Union[bool, str]] = None 'Enable thinking mode (for thinking models).' logprobs: Optional[bool] = None diff --git a/tests/test_client.py b/tests/test_client.py index 7b7ab38e..c212a11d 100644 --- a/tests/test_client.py +++ b/tests/test_client.py @@ -1490,10 +1490,40 @@ async def test_async_client_context_manager(): def test_generate_think_annotation_matches_chat(): - # The `think` parameter accepts bool or the 'low'/'medium'/'high' string levels. - # Client.generate must keep the same annotation as Client.chat and - # AsyncClient.generate so passing a string level does not raise a false type - # error (regression guard for the sync generate overloads/implementation). + # The `think` parameter accepts bool or string thinking levels (e.g. 'low', 'medium', 'high', 'xhigh', 'max'). + # Client.generate must keep the same annotation as Client.chat, + # AsyncClient.chat, and AsyncClient.generate so passing a string level does not + # raise a false type error (regression guard for the sync generate overloads/implementation). expected = inspect.signature(Client.chat).parameters['think'].annotation assert inspect.signature(Client.generate).parameters['think'].annotation == expected + assert inspect.signature(AsyncClient.chat).parameters['think'].annotation == expected assert inspect.signature(AsyncClient.generate).parameters['think'].annotation == expected + + +def test_client_chat_with_think_level(httpserver: HTTPServer): + httpserver.expect_ordered_request( + '/api/chat', + method='POST', + json={ + 'model': 'qwen3.8:27b', + 'messages': [{'role': 'user', 'content': 'Hello'}], + 'tools': [], + 'stream': False, + 'think': 'xhigh', + }, + ).respond_with_json( + { + 'model': 'qwen3.8:27b', + 'message': { + 'role': 'assistant', + 'content': 'Hi there.', + 'thinking': 'Thinking deeply...', + }, + } + ) + + client = Client(httpserver.url_for('/')) + response = client.chat('qwen3.8:27b', messages=[{'role': 'user', 'content': 'Hello'}], think='xhigh') + assert response['model'] == 'qwen3.8:27b' + assert response['message']['content'] == 'Hi there.' + assert response['message']['thinking'] == 'Thinking deeply...' diff --git a/tests/test_type_serialization.py b/tests/test_type_serialization.py index 02a69e95..e1af3512 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, GenerateRequest, Image def test_image_serialization_bytes(): @@ -105,3 +105,20 @@ 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'] + + +@pytest.mark.parametrize('level', ['low', 'medium', 'high', 'xhigh', 'max']) +def test_think_model_defined_levels_serialization(level): + chat_req = ChatRequest(model='test-model', messages=[{'role': 'user', 'content': 'hi'}], think=level) + assert chat_req.think == level + assert chat_req.model_dump(exclude_none=True)['think'] == level + + gen_req = GenerateRequest(model='test-model', think=level) + assert gen_req.think == level + assert gen_req.model_dump(exclude_none=True)['think'] == level + + +def test_think_boolean_serialization(): + assert ChatRequest(model='test-model', think=True).model_dump(exclude_none=True)['think'] is True + assert ChatRequest(model='test-model', think=False).model_dump(exclude_none=True)['think'] is False + assert 'think' not in ChatRequest(model='test-model', think=None).model_dump(exclude_none=True)