mirror of
https://github.com/ollama/ollama-python.git
synced 2026-09-30 22:27:23 +00:00
types/client: relax think type to support model-defined thinking levels (#744)
This commit is contained in:
+12
-12
@@ -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
@@ -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
@@ -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...'
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user