Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
24 changes: 12 additions & 12 deletions ollama/_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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,
Expand All @@ -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,
Expand Down Expand Up @@ -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,
Expand All @@ -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,
Expand All @@ -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,
Expand Down Expand Up @@ -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,
Expand All @@ -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,
Expand All @@ -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,
Expand Down Expand Up @@ -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,
Expand All @@ -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,
Expand All @@ -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,
Expand Down
4 changes: 2 additions & 2 deletions ollama/_types.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
38 changes: 34 additions & 4 deletions tests/test_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -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...'
19 changes: 18 additions & 1 deletion tests/test_type_serialization.py
Original file line number Diff line number Diff line change
Expand Up @@ -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():
Expand Down Expand Up @@ -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)
Loading