diff --git a/ollama/_client.py b/ollama/_client.py index 13dd6441..e768c604 100644 --- a/ollama/_client.py +++ b/ollama/_client.py @@ -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 @@ -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, @@ -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, @@ -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, @@ -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: ... @@ -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]: ... @@ -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]]: @@ -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, @@ -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, @@ -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, @@ -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: ... @@ -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]: ... @@ -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]]: diff --git a/ollama/_types.py b/ollama/_types.py index cf7264b5..942a0538 100644 --- a/ollama/_types.py +++ b/ollama/_types.py @@ -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 @@ -10,6 +11,8 @@ ByteSize, ConfigDict, Field, + TypeAdapter, + field_validator, model_serializer, ) from pydantic.json_schema import JsonSchemaValue @@ -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.' diff --git a/tests/test_client.py b/tests/test_client.py index c212a11d..81a0bee3 100644 --- a/tests/test_client.py +++ b/tests/test_client.py @@ -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 @@ -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'