types/client: relax think type to support model-defined thinking levels (#744)

This commit is contained in:
Juan Ezquerro LLanes
2026-09-28 18:15:40 -07:00
committed by GitHub
parent ed18cea18d
commit 8785556559
4 changed files with 66 additions and 19 deletions
+12 -12
View File
@@ -213,7 +213,7 @@ class Client(BaseClient):
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,
@@ -237,7 +237,7 @@ class Client(BaseClient):
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,
@@ -260,7 +260,7 @@ class Client(BaseClient):
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,
@@ -317,7 +317,7 @@ class Client(BaseClient):
*,
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,
@@ -333,7 +333,7 @@ class Client(BaseClient):
*,
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,
@@ -348,7 +348,7 @@ class Client(BaseClient):
*,
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,
@@ -869,7 +869,7 @@ class AsyncClient(BaseClient):
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,
@@ -893,7 +893,7 @@ class AsyncClient(BaseClient):
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,
@@ -916,7 +916,7 @@ class AsyncClient(BaseClient):
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,
@@ -972,7 +972,7 @@ class AsyncClient(BaseClient):
*,
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,
@@ -988,7 +988,7 @@ class AsyncClient(BaseClient):
*,
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,
@@ -1003,7 +1003,7 @@ class AsyncClient(BaseClient):
*,
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,
+2 -2
View File
@@ -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 @@ class ChatRequest(BaseGenerateRequest):
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
+34 -4
View File
@@ -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...'
+18 -1
View File
@@ -4,7 +4,7 @@ from pathlib import Path
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)