Skip to content
Closed
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
27 changes: 15 additions & 12 deletions google/genai/_automatic_function_calling_util.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,21 +14,18 @@
#

import inspect
import sys
import types as builtin_types
import typing
from typing import _GenericAlias, Any, Callable, get_args, get_origin, Literal, Optional, Union # type: ignore[attr-defined]
from typing import Any, Literal, Optional, Union, _GenericAlias, get_args, get_origin # type: ignore[attr-defined]

import pydantic

from . import _extra_utils
from . import types


if sys.version_info >= (3, 10):
VersionedUnionType = builtin_types.UnionType
else:
VersionedUnionType = typing._UnionGenericAlias # type: ignore[attr-defined]
_UNION_CLASSES = (builtin_types.UnionType, typing._UnionGenericAlias) # type: ignore[attr-defined]
_UNION_ORIGINS = (typing.Union, builtin_types.UnionType)


__all__ = [
Expand Down Expand Up @@ -102,10 +99,10 @@ def _is_default_value_compatible(
if (
isinstance(annotation, _GenericAlias)
or isinstance(annotation, builtin_types.GenericAlias)
or isinstance(annotation, VersionedUnionType)
or isinstance(annotation, _UNION_CLASSES)
):
origin = get_origin(annotation)
if origin in (Union, VersionedUnionType): # type: ignore[comparison-overlap]
if origin in _UNION_ORIGINS: # type: ignore[comparison-overlap]
return any(
_is_default_value_compatible(default_value, arg)
for arg in get_args(annotation)
Expand Down Expand Up @@ -160,7 +157,7 @@ def _parse_schema_from_parameter( # type: ignore[return]
schema.type = _py_builtin_type_to_schema_type[param.annotation]
return schema
if (
isinstance(param.annotation, VersionedUnionType)
isinstance(param.annotation, _UNION_CLASSES)
# only parse simple UnionType, example int | str | float | bool
# complex UnionType will be invoked in raise branch
and all(
Expand Down Expand Up @@ -199,8 +196,10 @@ def _parse_schema_from_parameter( # type: ignore[return]
raise ValueError(default_value_error_msg)
schema.default = param.default
return schema
if isinstance(param.annotation, _GenericAlias) or isinstance(
param.annotation, builtin_types.GenericAlias
if (
isinstance(param.annotation, _GenericAlias)
or isinstance(param.annotation, builtin_types.GenericAlias)
or isinstance(param.annotation, _UNION_CLASSES)
):
origin = get_origin(param.annotation)
args = get_args(param.annotation)
Expand Down Expand Up @@ -239,7 +238,11 @@ def _parse_schema_from_parameter( # type: ignore[return]
raise ValueError(default_value_error_msg)
schema.default = param.default
return schema
if origin is Union:
if origin in _UNION_ORIGINS:
if any(_extra_utils.is_annotation_pydantic_model(arg) for arg in args):
raise ValueError(
'Union containing Pydantic model is not supported in direct parsing'
)
schema.any_of = []
schema.type = _py_builtin_type_to_schema_type[dict]
unique_types = set()
Expand Down
3 changes: 3 additions & 0 deletions google/genai/types.py
Original file line number Diff line number Diff line change
Expand Up @@ -3287,6 +3287,9 @@ def convert_json_schema(
and field_name != 'additional_properties'
):
setattr(schema, field_name, field_value)
if schema.any_of and not schema.type:
if all(s.type == 'OBJECT' for s in schema.any_of):
schema.type = Type('OBJECT')

if (
schema.type == 'ARRAY'
Expand Down
Loading