Skip to content
Open
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
25 changes: 13 additions & 12 deletions ollama/_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@
import platform
import sys
import urllib.parse
from enum import Enum
from hashlib import sha256
from os import PathLike
from pathlib import Path
Expand Down Expand Up @@ -217,7 +218,7 @@ def generate(
logprobs: Optional[bool] = None,
top_logprobs: Optional[int] = None,
raw: bool = False,
format: Optional[Union[Literal['', 'json'], JsonSchemaValue]] = None,
format: Optional[Union[Literal['', 'json'], JsonSchemaValue, Type[Enum]]] = None,
images: Optional[Sequence[Union[str, bytes, Image]]] = None,
options: Optional[Union[Mapping[str, Any], Options]] = None,
keep_alive: Optional[Union[float, str]] = None,
Expand All @@ -241,7 +242,7 @@ def generate(
logprobs: Optional[bool] = None,
top_logprobs: Optional[int] = None,
raw: bool = False,
format: Optional[Union[Literal['', 'json'], JsonSchemaValue]] = None,
format: Optional[Union[Literal['', 'json'], JsonSchemaValue, Type[Enum]]] = None,
images: Optional[Sequence[Union[str, bytes, Image]]] = None,
options: Optional[Union[Mapping[str, Any], Options]] = None,
keep_alive: Optional[Union[float, str]] = None,
Expand All @@ -264,7 +265,7 @@ def generate(
logprobs: Optional[bool] = None,
top_logprobs: Optional[int] = None,
raw: Optional[bool] = None,
format: Optional[Union[Literal['', 'json'], JsonSchemaValue]] = None,
format: Optional[Union[Literal['', 'json'], JsonSchemaValue, Type[Enum]]] = None,
images: Optional[Sequence[Union[str, bytes, Image]]] = None,
options: Optional[Union[Mapping[str, Any], Options]] = None,
keep_alive: Optional[Union[float, str]] = None,
Expand Down Expand Up @@ -320,7 +321,7 @@ def chat(
think: Optional[Union[bool, str]] = None,
logprobs: Optional[bool] = None,
top_logprobs: Optional[int] = None,
format: Optional[Union[Literal['', 'json'], JsonSchemaValue]] = None,
format: Optional[Union[Literal['', 'json'], JsonSchemaValue, Type[Enum]]] = None,
options: Optional[Union[Mapping[str, Any], Options]] = None,
keep_alive: Optional[Union[float, str]] = None,
) -> ChatResponse: ...
Expand All @@ -336,7 +337,7 @@ def chat(
think: Optional[Union[bool, str]] = None,
logprobs: Optional[bool] = None,
top_logprobs: Optional[int] = None,
format: Optional[Union[Literal['', 'json'], JsonSchemaValue]] = None,
format: Optional[Union[Literal['', 'json'], JsonSchemaValue, Type[Enum]]] = None,
options: Optional[Union[Mapping[str, Any], Options]] = None,
keep_alive: Optional[Union[float, str]] = None,
) -> Iterator[ChatResponse]: ...
Expand All @@ -351,7 +352,7 @@ def chat(
think: Optional[Union[bool, str]] = None,
logprobs: Optional[bool] = None,
top_logprobs: Optional[int] = None,
format: Optional[Union[Literal['', 'json'], JsonSchemaValue]] = None,
format: Optional[Union[Literal['', 'json'], JsonSchemaValue, Type[Enum]]] = None,
options: Optional[Union[Mapping[str, Any], Options]] = None,
keep_alive: Optional[Union[float, str]] = None,
) -> Union[ChatResponse, Iterator[ChatResponse]]:
Expand Down Expand Up @@ -873,7 +874,7 @@ async def generate(
logprobs: Optional[bool] = None,
top_logprobs: Optional[int] = None,
raw: bool = False,
format: Optional[Union[Literal['', 'json'], JsonSchemaValue]] = None,
format: Optional[Union[Literal['', 'json'], JsonSchemaValue, Type[Enum]]] = None,
images: Optional[Sequence[Union[str, bytes, Image]]] = None,
options: Optional[Union[Mapping[str, Any], Options]] = None,
keep_alive: Optional[Union[float, str]] = None,
Expand All @@ -897,7 +898,7 @@ async def generate(
logprobs: Optional[bool] = None,
top_logprobs: Optional[int] = None,
raw: bool = False,
format: Optional[Union[Literal['', 'json'], JsonSchemaValue]] = None,
format: Optional[Union[Literal['', 'json'], JsonSchemaValue, Type[Enum]]] = None,
images: Optional[Sequence[Union[str, bytes, Image]]] = None,
options: Optional[Union[Mapping[str, Any], Options]] = None,
keep_alive: Optional[Union[float, str]] = None,
Expand All @@ -920,7 +921,7 @@ async def generate(
logprobs: Optional[bool] = None,
top_logprobs: Optional[int] = None,
raw: Optional[bool] = None,
format: Optional[Union[Literal['', 'json'], JsonSchemaValue]] = None,
format: Optional[Union[Literal['', 'json'], JsonSchemaValue, Type[Enum]]] = None,
images: Optional[Sequence[Union[str, bytes, Image]]] = None,
options: Optional[Union[Mapping[str, Any], Options]] = None,
keep_alive: Optional[Union[float, str]] = None,
Expand Down Expand Up @@ -975,7 +976,7 @@ async def chat(
think: Optional[Union[bool, str]] = None,
logprobs: Optional[bool] = None,
top_logprobs: Optional[int] = None,
format: Optional[Union[Literal['', 'json'], JsonSchemaValue]] = None,
format: Optional[Union[Literal['', 'json'], JsonSchemaValue, Type[Enum]]] = None,
options: Optional[Union[Mapping[str, Any], Options]] = None,
keep_alive: Optional[Union[float, str]] = None,
) -> ChatResponse: ...
Expand All @@ -991,7 +992,7 @@ async def chat(
think: Optional[Union[bool, str]] = None,
logprobs: Optional[bool] = None,
top_logprobs: Optional[int] = None,
format: Optional[Union[Literal['', 'json'], JsonSchemaValue]] = None,
format: Optional[Union[Literal['', 'json'], JsonSchemaValue, Type[Enum]]] = None,
options: Optional[Union[Mapping[str, Any], Options]] = None,
keep_alive: Optional[Union[float, str]] = None,
) -> AsyncIterator[ChatResponse]: ...
Expand All @@ -1006,7 +1007,7 @@ async def chat(
think: Optional[Union[bool, str]] = None,
logprobs: Optional[bool] = None,
top_logprobs: Optional[int] = None,
format: Optional[Union[Literal['', 'json'], JsonSchemaValue]] = None,
format: Optional[Union[Literal['', 'json'], JsonSchemaValue, Type[Enum]]] = None,
options: Optional[Union[Mapping[str, Any], Options]] = None,
keep_alive: Optional[Union[float, str]] = None,
) -> Union[ChatResponse, AsyncIterator[ChatResponse]]:
Expand Down
10 changes: 10 additions & 0 deletions ollama/_types.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@
import json
from base64 import b64decode, b64encode
from datetime import datetime
from enum import Enum
from pathlib import Path
from typing import Any, Dict, List, Mapping, Optional, Sequence, Union

Expand All @@ -10,6 +11,8 @@
ByteSize,
ConfigDict,
Field,
TypeAdapter,
field_validator,
model_serializer,
)
from pydantic.json_schema import JsonSchemaValue
Expand Down Expand Up @@ -154,6 +157,13 @@ class BaseGenerateRequest(BaseStreamableRequest):
format: Optional[Union[Literal['', 'json'], JsonSchemaValue]] = None
'Format of the response.'

@field_validator('format', mode='before')
@classmethod
def _enum_format(cls, value: Any) -> Any:
if isinstance(value, type) and issubclass(value, Enum):
return TypeAdapter(value).json_schema()
return value

keep_alive: Optional[Union[float, str]] = None
'Keep model alive for the specified duration.'

Expand Down
32 changes: 32 additions & 0 deletions tests/test_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,10 +4,13 @@
import os
import re
import tempfile
from enum import Enum
from pathlib import Path
from typing import Any

import pytest
from httpx import MockTransport
from httpx import Request as httpxRequest
from httpx import Response as httpxResponse
from pydantic import BaseModel
from pytest_httpserver import HTTPServer, URIPattern
Expand All @@ -22,6 +25,35 @@
pytestmark = pytest.mark.anyio


class Answer(Enum):
yes = 'yes'
no = 'no'


@pytest.mark.parametrize('method', ['chat', 'generate'])
def test_client_accepts_enum_response_format(method: str):
def respond(request: httpxRequest) -> httpxResponse:
assert request.url.path == f'/api/{method}'
assert json.loads(request.content)['format'] == {'enum': ['yes', 'no'], 'title': 'Answer', 'type': 'string'}
return httpxResponse(200, json={'model': 'dummy', 'message': {'role': 'assistant', 'content': '"yes"'}} if method == 'chat' else {'model': 'dummy', 'response': '"yes"'})

with Client(transport=MockTransport(respond)) as client:
response = getattr(client, method)(model='dummy', format=Answer)
assert response.model == 'dummy'


@pytest.mark.parametrize('method', ['chat', 'generate'])
async def test_async_client_accepts_enum_response_format(method: str):
def respond(request: httpxRequest) -> httpxResponse:
assert request.url.path == f'/api/{method}'
assert json.loads(request.content)['format'] == {'enum': ['yes', 'no'], 'title': 'Answer', 'type': 'string'}
return httpxResponse(200, json={'model': 'dummy', 'message': {'role': 'assistant', 'content': '"yes"'}} if method == 'chat' else {'model': 'dummy', 'response': '"yes"'})

async with AsyncClient(transport=MockTransport(respond)) as client:
response = await getattr(client, method)(model='dummy', format=Answer)
assert response.model == 'dummy'


@pytest.fixture
def anyio_backend():
return 'asyncio'
Expand Down