Skip to content
Merged
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: 5 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -40,6 +40,11 @@ to include examples, links to docs, or any other relevant information.

### Fixed

- `contrib.google_adk_agents`: agents with an `output_schema` no longer fail every workflow task
when calling the model. The schema type is now sent to the model activity as its JSON schema.
Custom Pydantic schema generation is preserved.
Integer-valued output enums are normalized to strings to match Google GenAI.

- `temporalio.contrib.strands` activity and MCP tools now give the model the
Activity's failure message and expose its exception to after-tool hooks.

Expand Down
36 changes: 36 additions & 0 deletions temporalio/contrib/google_adk_agents/_model.py
Original file line number Diff line number Diff line change
@@ -1,10 +1,13 @@
from collections.abc import AsyncGenerator, Callable
from dataclasses import dataclass
from datetime import timedelta
from enum import Enum

from google.adk.models import BaseLlm, LLMRegistry
from google.adk.models.llm_request import LlmRequest
from google.adk.models.llm_response import LlmResponse
from google.genai import types
from pydantic import BaseModel, TypeAdapter

import temporalio.workflow
from temporalio import activity, workflow
Expand Down Expand Up @@ -91,6 +94,38 @@ async def invoke_model_streaming(
return responses


def _with_serializable_response_schema(llm_request: LlmRequest) -> LlmRequest:
"""Return the request with a ``response_schema`` that can be serialized.

ADK stores an agent's ``output_schema`` on the request as a Python type
(for example a Pydantic model class), which the payload converter cannot
serialize. google-genai and ADK's LiteLlm both turn such a type into its
JSON schema before calling the model, so sending the JSON schema instead
is equivalent. Pydantic model classes use their ``model_json_schema``
method to preserve custom schema generation. Integer-valued enums are
normalized to string enums to match google-genai's enum handling.
"""
schema = llm_request.config.response_schema
if schema is None or isinstance(schema, (dict, types.Schema)):
return llm_request
if isinstance(schema, type) and issubclass(schema, BaseModel):
response_schema = schema.model_json_schema()
else:
response_schema = TypeAdapter(schema).json_schema()
if (
isinstance(schema, type)
and issubclass(schema, Enum)
and any(isinstance(member.value, int) for member in schema)
):
response_schema["type"] = "string"
response_schema["enum"] = [str(member.value) for member in schema]
request = llm_request.model_copy()
request.config = llm_request.config.model_copy(
update={"response_schema": response_schema}
)
return request


class TemporalModel(BaseLlm):
"""A Temporal-based LLM model that executes model invocations as activities."""

Expand Down Expand Up @@ -183,6 +218,7 @@ async def generate_content_async(
if agent_name:
config["summary"] = agent_name

llm_request = _with_serializable_response_schema(llm_request)
if stream:
if self._streaming_topic is None:
raise ApplicationError(
Expand Down
185 changes: 185 additions & 0 deletions tests/contrib/google_adk_agents/test_google_adk_agents.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,7 @@
from abc import ABC, abstractmethod
from collections.abc import AsyncGenerator
from datetime import timedelta
from enum import Enum, IntEnum
from typing import Any

import pytest
Expand All @@ -43,6 +44,7 @@
from opentelemetry.sdk.trace.export import SimpleSpanProcessor
from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter
from opentelemetry.trace import get_tracer_provider, set_tracer_provider
from pydantic import BaseModel, TypeAdapter

import temporalio.contrib.google_adk_agents.workflow
from temporalio import activity, workflow
Expand All @@ -53,6 +55,9 @@
TemporalMcpToolSetProvider,
TemporalModel,
)
from temporalio.contrib.google_adk_agents._model import (
_with_serializable_response_schema,
)
from temporalio.contrib.opentelemetry import OpenTelemetryPlugin, create_tracer_provider
from temporalio.worker import Worker
from temporalio.workflow import ActivityConfig
Expand Down Expand Up @@ -1168,3 +1173,183 @@ async def my_activity(city: str, count: int = 1) -> str:
assert params == ["city", "count"]
assert sig.parameters["city"].annotation is str
assert sig.parameters["count"].default == 1


class CityWeather(BaseModel):
city: str
temperature_c: float


class NumericChoice(IntEnum):
FIRST = 10
SECOND = 20


class IntegerChoice(Enum):
FIRST = 1
SECOND = 2


class MixedChoice(Enum):
FIRST = 1
SECOND = "other"


class StringChoice(str, Enum):
FIRST = "first"
SECOND = "second"


class OutputSchemaModel(TestModel):
def responses(self) -> list[LlmResponse]:
return [
LlmResponse(
content=Content(
role="model",
parts=[Part(text='{"city": "Paris", "temperature_c": 17.5}')],
)
)
]

@classmethod
def supported_models(cls) -> list[str]:
return ["output_schema_model"]


@workflow.defn
class OutputSchemaAgentWorkflow:
@workflow.run
async def run(self, prompt: str) -> dict[str, Any] | None:
agent = LlmAgent(
name="output_schema_agent",
model=TemporalModel("output_schema_model"),
output_schema=CityWeather,
output_key="weather",
)
runner = InMemoryRunner(agent=agent, app_name="output_schema_app")
session = await runner.session_service.create_session(
app_name="output_schema_app", user_id="test"
)
async with Aclosing(
runner.run_async(
user_id="test",
session_id=session.id,
new_message=types.Content(role="user", parts=[types.Part(text=prompt)]),
)
) as agen:
async for _ in agen:
pass

final_session = await runner.session_service.get_session(
app_name="output_schema_app", user_id="test", session_id=session.id
)
return final_session.state.get("weather") if final_session else None


@pytest.mark.asyncio
async def test_agent_with_output_schema(client: Client):
LLMRegistry.register(OutputSchemaModel)

new_config = client.config()
new_config["plugins"] = [GoogleAdkPlugin()]
client = Client(**new_config)

async with Worker(
client,
task_queue="adk-task-queue-output-schema",
workflows=[OutputSchemaAgentWorkflow],
max_cached_workflows=0,
):
result = await client.execute_workflow(
OutputSchemaAgentWorkflow.run,
"What is the weather in Paris?",
id=f"output-schema-agent-workflow-{uuid.uuid4()}",
task_queue="adk-task-queue-output-schema",
execution_timeout=timedelta(seconds=60),
)

assert result == {"city": "Paris", "temperature_c": 17.5}


@pytest.mark.parametrize("schema", [CityWeather, list[CityWeather], StringChoice])
def test_output_schema_type_sent_as_json_schema(schema: Any) -> None:
request = LlmRequest(
model="gemini-2.0-flash",
contents=[Content(role="user", parts=[Part(text="hello")])],
config=types.GenerateContentConfig(),
)
request.set_output_schema(schema)

converted = _with_serializable_response_schema(request)

assert request.config.response_schema is schema
assert converted.config.response_mime_type == "application/json"
converter = GoogleAdkPlugin()._configure_data_converter(None)
payloads = converter.payload_converter.to_payloads([converted])
serialized = json.loads(payloads[0].data)
assert serialized["config"]["response_schema"] == TypeAdapter(schema).json_schema()


def test_output_schema_preserves_custom_model_json_schema() -> None:
class CustomCityWeather(CityWeather):
@classmethod
def model_json_schema(cls, *args: Any, **kwargs: Any) -> dict[str, Any]:
"""Include the cities supported by the weather model."""
schema = super().model_json_schema(*args, **kwargs)
schema["properties"]["city"]["enum"] = ["Paris", "London"]
return schema

request = LlmRequest(
model="gemini-2.0-flash",
config=types.GenerateContentConfig(),
)
request.set_output_schema(CustomCityWeather)

converted = _with_serializable_response_schema(request)
converter = GoogleAdkPlugin()._configure_data_converter(None).payload_converter
payloads = converter.to_payloads([converted])
restored = converter.from_payloads(payloads, [LlmRequest])[0]
response_schema = types.Schema.model_validate(restored.config.response_schema)

assert request.config.response_schema is CustomCityWeather
assert response_schema.properties is not None
assert response_schema.properties["city"].enum == ["Paris", "London"]


@pytest.mark.parametrize(
("schema", "expected_values"),
[
(NumericChoice, ["10", "20"]),
(IntegerChoice, ["1", "2"]),
(MixedChoice, ["1", "other"]),
],
)
def test_output_schema_integer_enum_is_serializable(
schema: type[Enum], expected_values: list[str]
) -> None:
request = LlmRequest(
model="gemini-2.0-flash",
config=types.GenerateContentConfig(),
)
request.set_output_schema(schema)

converted = _with_serializable_response_schema(request)

assert request.config.response_schema is schema
converter = GoogleAdkPlugin()._configure_data_converter(None).payload_converter
payloads = converter.to_payloads([converted])
restored = converter.from_payloads(payloads, [LlmRequest])[0]
response_schema = types.Schema.model_validate(restored.config.response_schema)
assert response_schema.type == types.Type.STRING
assert response_schema.enum == expected_values


def test_json_output_schema_left_unchanged() -> None:
request = LlmRequest(
model="gemini-2.0-flash",
contents=[Content(role="user", parts=[Part(text="hello")])],
config=types.GenerateContentConfig(),
)
request.set_output_schema(CityWeather.model_json_schema())

assert _with_serializable_response_schema(request) is request
Loading