mirror of
https://github.com/encounter/adk-python.git
synced 2026-07-09 18:19:28 -07:00
fix: Support Generator and Async Generator tool declaration in JSON schema
Co-authored-by: Xiang (Sean) Zhou <seanzhougoogle@google.com> PiperOrigin-RevId: 856713741
This commit is contained in:
committed by
Copybara-Service
parent
ed2c3ebde9
commit
19555e7dce
@@ -24,10 +24,13 @@ allowing us to delegate schema generation complexity to Pydantic.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import collections.abc
|
||||
import inspect
|
||||
import logging
|
||||
from typing import Any
|
||||
from typing import Callable
|
||||
from typing import get_args
|
||||
from typing import get_origin
|
||||
from typing import get_type_hints
|
||||
from typing import Optional
|
||||
from typing import Type
|
||||
@@ -145,6 +148,19 @@ def _build_response_json_schema(
|
||||
except TypeError:
|
||||
pass
|
||||
|
||||
# Handle AsyncGenerator and Generator return types (streaming tools)
|
||||
# AsyncGenerator[YieldType, SendType] -> use YieldType as response schema
|
||||
# Generator[YieldType, SendType, ReturnType] -> use YieldType as response schema
|
||||
origin = get_origin(return_annotation)
|
||||
if origin is not None and (
|
||||
origin is collections.abc.AsyncGenerator
|
||||
or origin is collections.abc.Generator
|
||||
):
|
||||
type_args = get_args(return_annotation)
|
||||
if type_args:
|
||||
# First type argument is the yield type
|
||||
return_annotation = type_args[0]
|
||||
|
||||
try:
|
||||
adapter = pydantic.TypeAdapter(
|
||||
return_annotation,
|
||||
|
||||
@@ -23,6 +23,8 @@ from __future__ import annotations
|
||||
from collections.abc import Sequence
|
||||
from enum import Enum
|
||||
from typing import Any
|
||||
from typing import AsyncGenerator
|
||||
from typing import Generator
|
||||
from typing import Literal
|
||||
from typing import Optional
|
||||
|
||||
@@ -840,3 +842,81 @@ class TestPydanticModelAsFunction(parameterized.TestCase):
|
||||
# When passing a BaseModel, there is no function return, so response schema
|
||||
# is None
|
||||
self.assertIsNone(decl.response_json_schema)
|
||||
|
||||
|
||||
class TestStreamingReturnTypes(parameterized.TestCase):
|
||||
"""Tests for AsyncGenerator and Generator return types (streaming tools)."""
|
||||
|
||||
def test_async_generator_string_yield(self):
|
||||
"""Test AsyncGenerator[str, None] return type extracts str as response."""
|
||||
|
||||
async def streaming_tool(param: str) -> AsyncGenerator[str, None]:
|
||||
"""A streaming tool that yields strings."""
|
||||
yield param
|
||||
|
||||
decl = build_function_declaration_with_json_schema(streaming_tool)
|
||||
|
||||
self.assertEqual(decl.name, "streaming_tool")
|
||||
self.assertIsNotNone(decl.parameters_json_schema)
|
||||
self.assertEqual(
|
||||
decl.parameters_json_schema["properties"]["param"]["type"], "string"
|
||||
)
|
||||
# Should extract str from AsyncGenerator[str, None]
|
||||
self.assertEqual(decl.response_json_schema, {"type": "string"})
|
||||
|
||||
def test_async_generator_int_yield(self):
|
||||
"""Test AsyncGenerator[int, None] return type extracts int as response."""
|
||||
|
||||
async def counter(start: int) -> AsyncGenerator[int, None]:
|
||||
"""A streaming counter."""
|
||||
yield start
|
||||
|
||||
decl = build_function_declaration_with_json_schema(counter)
|
||||
|
||||
self.assertEqual(decl.name, "counter")
|
||||
# Should extract int from AsyncGenerator[int, None]
|
||||
self.assertEqual(decl.response_json_schema, {"type": "integer"})
|
||||
|
||||
def test_async_generator_dict_yield(self):
|
||||
"""Test AsyncGenerator[dict[str, str], None] return type."""
|
||||
|
||||
async def streaming_dict(
|
||||
param: str,
|
||||
) -> AsyncGenerator[dict[str, str], None]:
|
||||
"""A streaming tool that yields dicts."""
|
||||
yield {"result": param}
|
||||
|
||||
decl = build_function_declaration_with_json_schema(streaming_dict)
|
||||
|
||||
self.assertEqual(decl.name, "streaming_dict")
|
||||
# Should extract dict[str, str] from AsyncGenerator
|
||||
self.assertEqual(
|
||||
decl.response_json_schema,
|
||||
{"additionalProperties": {"type": "string"}, "type": "object"},
|
||||
)
|
||||
|
||||
def test_generator_string_yield(self):
|
||||
"""Test Generator[str, None, None] return type extracts str as response."""
|
||||
|
||||
def sync_streaming_tool(param: str) -> Generator[str, None, None]:
|
||||
"""A sync streaming tool that yields strings."""
|
||||
yield param
|
||||
|
||||
decl = build_function_declaration_with_json_schema(sync_streaming_tool)
|
||||
|
||||
self.assertEqual(decl.name, "sync_streaming_tool")
|
||||
# Should extract str from Generator[str, None, None]
|
||||
self.assertEqual(decl.response_json_schema, {"type": "string"})
|
||||
|
||||
def test_generator_int_yield(self):
|
||||
"""Test Generator[int, None, None] return type extracts int as response."""
|
||||
|
||||
def sync_counter(start: int) -> Generator[int, None, None]:
|
||||
"""A sync streaming counter."""
|
||||
yield start
|
||||
|
||||
decl = build_function_declaration_with_json_schema(sync_counter)
|
||||
|
||||
self.assertEqual(decl.name, "sync_counter")
|
||||
# Should extract int from Generator[int, None, None]
|
||||
self.assertEqual(decl.response_json_schema, {"type": "integer"})
|
||||
|
||||
Reference in New Issue
Block a user