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
21 changes: 20 additions & 1 deletion ollama/_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
from typing import Callable, Union

import pydantic
from typing_extensions import get_type_hints

from ollama._types import Tool

Expand Down Expand Up @@ -53,14 +54,32 @@ def _parse_docstring(doc_string: Union[str, None]) -> dict[str, str]:
return parsed_docstring


def _resolve_type_hints(func: Callable) -> dict:
# String annotations (quoted, or under `from __future__ import annotations`) name types
# from the function's module, which pydantic cannot see from here, so evaluate them there.
try:
return get_type_hints(inspect.unwrap(func), include_extras=True)
except Exception:
return {}


def _parameter_annotation(name: str, parameter: inspect.Parameter, type_hints: dict):
if parameter.annotation is inspect.Parameter.empty:
return str
if isinstance(parameter.annotation, str):
return type_hints.get(name, parameter.annotation)
return parameter.annotation


def convert_function_to_tool(func: Callable) -> Tool:
doc_string_hash = str(hash(inspect.getdoc(func)))
parsed_docstring = _parse_docstring(inspect.getdoc(func))
type_hints = _resolve_type_hints(func)
schema = type(
func.__name__,
(pydantic.BaseModel,),
{
'__annotations__': {k: v.annotation if v.annotation != inspect._empty else str for k, v in inspect.signature(func).parameters.items()},
'__annotations__': {k: _parameter_annotation(k, v, type_hints) for k, v in inspect.signature(func).parameters.items()},
'__signature__': inspect.signature(func),
'__doc__': parsed_docstring[doc_string_hash],
},
Expand Down
20 changes: 19 additions & 1 deletion tests/test_utils.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
import json
import sys
from typing import Dict, List, Mapping, Sequence, Set, Tuple, Union
from typing import Dict, List, Literal, Mapping, Optional, Sequence, Set, Tuple, Union

from ollama._utils import convert_function_to_tool

Expand Down Expand Up @@ -256,3 +256,21 @@ def func_with_parentheses_and_args(a: int, b: int):
tool = convert_function_to_tool(func_with_parentheses_and_args).model_dump()
assert tool['function']['parameters']['properties']['a']['description'] == 'First (:thing) number to add'
assert tool['function']['parameters']['properties']['b']['description'] == 'Second number to add'


def test_function_with_string_annotations():
# Annotations are strings under `from __future__ import annotations`, or when quoted.
def get_weather(city: 'str', unit: "Literal['celsius', 'fahrenheit']", days: 'Optional[int]' = None) -> 'str':
"""
Get the weather for a city.
Args:
city: The city
unit: The unit
days: Number of days
"""

tool = convert_function_to_tool(get_weather).model_dump()
assert tool['function']['parameters']['properties']['city']['type'] == 'string'
assert tool['function']['parameters']['properties']['unit']['type'] == 'string'
assert tool['function']['parameters']['properties']['days']['type'] == 'integer'
assert tool['function']['parameters']['required'] == ['city', 'unit']