types: relax type for tools (#550)
test / test (push) Has been cancelled
test / lint (push) Has been cancelled

This commit is contained in:
Parth Sareen
2025-08-05 15:59:56 -07:00
committed by GitHub
parent dad9e1ca3a
commit 34e98bd237
3 changed files with 10 additions and 8 deletions
+2 -2
View File
@@ -79,7 +79,7 @@ class SubscriptableBaseModel(BaseModel):
if key in self.model_fields_set: if key in self.model_fields_set:
return True return True
if value := self.model_fields.get(key): if value := self.__class__.model_fields.get(key):
return value.default is not None return value.default is not None
return False return False
@@ -313,7 +313,7 @@ class Message(SubscriptableBaseModel):
class Tool(SubscriptableBaseModel): class Tool(SubscriptableBaseModel):
type: Optional[Literal['function']] = 'function' type: Optional[str] = 'function'
class Function(SubscriptableBaseModel): class Function(SubscriptableBaseModel):
name: Optional[str] = None name: Optional[str] = None
+2 -1
View File
@@ -79,11 +79,12 @@ def convert_function_to_tool(func: Callable) -> Tool:
} }
tool = Tool( tool = Tool(
type='function',
function=Tool.Function( function=Tool.Function(
name=func.__name__, name=func.__name__,
description=schema.get('description', ''), description=schema.get('description', ''),
parameters=Tool.Function.Parameters(**schema), parameters=Tool.Function.Parameters(**schema),
) ),
) )
return Tool.model_validate(tool) return Tool.model_validate(tool)
+6 -5
View File
@@ -8,7 +8,7 @@ from typing import Any
import pytest import pytest
from httpx import Response as httpxResponse from httpx import Response as httpxResponse
from pydantic import BaseModel, ValidationError from pydantic import BaseModel
from pytest_httpserver import HTTPServer, URIPattern from pytest_httpserver import HTTPServer, URIPattern
from werkzeug.wrappers import Request, Response from werkzeug.wrappers import Request, Response
@@ -1136,10 +1136,11 @@ def test_copy_tools():
def test_tool_validation(): def test_tool_validation():
# Raises ValidationError when used as it is a generator arbitrary_tool = {'type': 'custom_type', 'function': {'name': 'test'}}
with pytest.raises(ValidationError): tools = list(_copy_tools([arbitrary_tool]))
invalid_tool = {'type': 'invalid_type', 'function': {'name': 'test'}} assert len(tools) == 1
list(_copy_tools([invalid_tool])) assert tools[0].type == 'custom_type'
assert tools[0].function.name == 'test'
def test_client_connection_error(): def test_client_connection_error():