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
5 changes: 2 additions & 3 deletions src/google/adk/utils/_schema_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,6 +37,7 @@
from typing_extensions import Annotated

from . import _json_utils
from .content_utils import extract_text_from_content

logger = logging.getLogger("google_adk." + __name__)

Expand Down Expand Up @@ -312,9 +313,7 @@ def _validate_python_object(val: Any) -> Any:
return _validate_python_object(data)

if isinstance(data, types.Content):
# Extract text part
text_parts = [p.text for p in data.parts if p.text] if data.parts else []
text_str = "".join(text_parts)
text_str = extract_text_from_content(data)

# Validate the text
if schema is str:
Expand Down
12 changes: 6 additions & 6 deletions src/google/adk/workflow/_function_node.py
Original file line number Diff line number Diff line change
Expand Up @@ -44,6 +44,7 @@
from ..utils._schema_utils import annotation_accepts_content
from ..utils._schema_utils import annotation_expects_str
from ..utils._sync_runner import _SYNC_CALLABLE_RUNNER
from ..utils.content_utils import extract_text_from_content
from ..utils.context_utils import Aclosing
from ._base_node import BaseNode
from ._errors import WorkflowConfigurationError
Expand Down Expand Up @@ -205,20 +206,19 @@ def _step_in_pool() -> tuple[str, Any]:
def _content_to_str(
content: types.Content, func_name: str, param_name: str
) -> str:
"""Extracts text from a Content object, warning on non-text parts."""
texts = []
"""Extracts answer text from Content, warning on non-text parts."""
for part in content.parts or []:
if part.text is not None:
texts.append(part.text)
elif part.inline_data or part.file_data or part.executable_code:
if part.text is None and (
part.inline_data or part.file_data or part.executable_code
):
logger.warning(
'Parameter "%s" of function "%s" expects str but received'
" Content with non-text parts (e.g. inline_data, file_data)."
" Non-text parts are dropped during auto-conversion.",
param_name,
func_name,
)
return "".join(texts)
return extract_text_from_content(content)


_expects_str = annotation_expects_str
Expand Down
26 changes: 26 additions & 0 deletions tests/unittests/utils/test_schema_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -300,6 +300,32 @@ def test_strip_json_code_fence_variations(self):
class TestValidateNodeData:
"""Tests for validate_node_data function."""

@pytest.mark.parametrize("preserve_content", [False, True])
@pytest.mark.parametrize("schema", [str, SampleModel])
def test_content_thoughts_are_excluded_from_validated_text(
self, schema, preserve_content
):
"""Only answer text is validated or returned as a typed node payload."""
answer = "hello" if schema is str else '{"name": "test", "value": 42}'
data = types.Content(
role="model",
parts=[
types.Part(text="Let me think first.", thought=True),
types.Part(text=answer),
],
)

result = validate_node_data(schema, data, preserve_content=preserve_content)

if preserve_content:
assert result == types.Content(
role="model", parts=[types.Part(text=answer)]
)
else:
assert result == (
answer if schema is str else {"name": "test", "value": 42}
)

def test_none_schema_or_data_returns_data(self):
"""Bypasses validation if schema or data is None."""
assert validate_node_data(None, "some_data") == "some_data"
Expand Down
37 changes: 37 additions & 0 deletions tests/unittests/workflow/test_function_node.py
Original file line number Diff line number Diff line change
Expand Up @@ -1418,6 +1418,43 @@ def produce() -> dict:
assert received == [_OutputModel(name='test', value=42)]


@pytest.mark.parametrize('structured', [False, True])
async def test_typed_function_receives_answer_without_thoughts(structured):
"""A workflow passes the answer, rather than thoughts, into typed functions."""
received = []

def process_text(node_input: str) -> str:
received.append(node_input)
return 'ok'

def process_model(node_input: _OutputModel) -> str:
received.append(node_input)
return 'ok'

answer = '{"name": "test", "value": 42}' if structured else 'hello'

def produce() -> Event:
return Event(
output=types.Content(
role='model',
parts=[
types.Part(text='Let me think first.', thought=True),
types.Part(text=answer),
],
)
)

process = process_model if structured else process_text
agent = Workflow(
name='typed_answer', edges=[(START, produce), (produce, process)]
)
await run_workflow(agent)

assert received == (
[_OutputModel(name='test', value=42)] if structured else ['hello']
)


@pytest.mark.asyncio
async def test_input_schema_rejects_invalid_dict(
request: pytest.FixtureRequest,
Expand Down
Loading