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
195 changes: 195 additions & 0 deletions tests/agent/test_workflow_execution.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,17 +22,26 @@
from __future__ import annotations

import asyncio
from collections.abc import AsyncGenerator
from typing import Any

import pytest
from google.adk.agents import BaseAgent
from google.adk.agents.invocation_context import InvocationContext
from google.adk.events import Event, EventActions
from google.adk.models.base_llm import BaseLlm
from google.adk.models.llm_response import LlmResponse
from google.adk.runners import InMemoryRunner
from google.adk.sessions import InMemorySessionService, Session
from google.genai import types
from packaging.version import Version

import veadk.utils.patches as patches
from veadk import Agent
from veadk.agents.loop_agent import LoopAgent
from veadk.agents.parallel_agent import ParallelAgent
from veadk.agents.sequential_agent import SequentialAgent
from veadk.utils.adk_compat import get_adk_version


class _ExecutionTracker:
Expand Down Expand Up @@ -74,6 +83,40 @@ async def _run_async_impl(self, ctx: InvocationContext):
)


class _TransferLlm(BaseLlm):
async def generate_content_async(
self, llm_request, stream: bool = False
) -> AsyncGenerator[LlmResponse, None]:
del llm_request, stream
yield LlmResponse(
content=types.Content(
role="model",
parts=[
types.Part(
function_call=types.FunctionCall(
id="transfer-worker",
name="transfer_to_agent",
args={"agent_name": "worker"},
)
)
],
)
)


class _DirectLlm(BaseLlm):
async def generate_content_async(
self, llm_request, stream: bool = False
) -> AsyncGenerator[LlmResponse, None]:
del llm_request, stream
yield LlmResponse(
content=types.Content(
role="model",
parts=[types.Part(text="worker-ok")],
)
)


def _invocation_context(agent: BaseAgent) -> InvocationContext:
return InvocationContext(
session_service=InMemorySessionService(),
Expand Down Expand Up @@ -148,6 +191,158 @@ async def test_loop_agent_repeats_until_a_child_escalates() -> None:
assert events[-1].actions.escalate is True


def _workflow_peer(agent_type: str, tracker: _ExecutionTracker) -> BaseAgent:
leaf = _DeterministicAgent(
name=f"{agent_type}_peer_leaf",
tracker=tracker,
escalate_on_call=1 if agent_type == "loop" else None,
)
if agent_type == "parallel":
return ParallelAgent(name="transfer_peer", sub_agents=[leaf])
if agent_type == "sequential":
return SequentialAgent(name="transfer_peer", sub_agents=[leaf])
if agent_type == "loop":
return LoopAgent(name="transfer_peer", sub_agents=[leaf], max_iterations=2)
return _DeterministicAgent(name="transfer_peer", tracker=tracker)


def test_adk_2_0_to_2_1_workflow_mode_compat_does_not_change_serialization() -> None:
if not Version("2.0.0") <= get_adk_version() < Version("2.2.0"):
pytest.skip("The compatibility attribute is limited to ADK 2.0 and 2.1")

for workflow in (
ParallelAgent(name="parallel", sub_agents=[]),
SequentialAgent(name="sequential", sub_agents=[]),
LoopAgent(name="loop", sub_agents=[], max_iterations=1),
_DeterministicAgent(name="custom", tracker=_ExecutionTracker()),
):
assert workflow.mode is None
assert "mode" not in workflow.model_dump()


@pytest.mark.parametrize("adk_version", ("2.0.0", "2.1.0", "2.1.99"))
def test_workflow_mode_patch_covers_all_affected_adk_versions(
monkeypatch: pytest.MonkeyPatch,
adk_version: str,
) -> None:
import google.adk.agents as adk_agents

base_agent = type("BaseAgent", (), {})
workflow_types = [
type(name, (base_agent,), {})
for name in ("ParallelAgent", "SequentialAgent", "LoopAgent", "CustomAgent")
]
preconfigured_agent = type("PreconfiguredAgent", (base_agent,), {"mode": "task"})
monkeypatch.setattr(
patches,
"get_adk_version",
lambda: Version(adk_version),
)
monkeypatch.setattr(adk_agents, "BaseAgent", base_agent)

patches.patch_adk_workflow_agent_mode()
patches.patch_adk_workflow_agent_mode()

assert base_agent.mode is None
assert all(agent_type.mode is None for agent_type in workflow_types)
assert preconfigured_agent.mode == "task"


@pytest.mark.parametrize("adk_version", ("1.34.3", "2.2.0", "3.0.0"))
def test_workflow_mode_patch_does_not_touch_unaffected_adk_versions(
monkeypatch: pytest.MonkeyPatch,
adk_version: str,
) -> None:
import google.adk.agents as adk_agents

base_agent = type("BaseAgent", (), {})
agent_types = [
type(name, (base_agent,), {})
for name in ("ParallelAgent", "SequentialAgent", "LoopAgent", "CustomAgent")
]
monkeypatch.setattr(
patches,
"get_adk_version",
lambda: Version(adk_version),
)
monkeypatch.setattr(adk_agents, "BaseAgent", base_agent)

patches.patch_adk_workflow_agent_mode()

assert not hasattr(base_agent, "mode")
assert all(not hasattr(agent_type, "mode") for agent_type in agent_types)


@pytest.mark.parametrize("peer_type", ("parallel", "sequential", "loop", "custom"))
@pytest.mark.asyncio
async def test_runner_handles_deep_mixed_workflow_peer_transfer(
peer_type: str,
) -> None:
"""Delegated LLMs support every workflow peer affected before ADK 2.2."""
tracker = _ExecutionTracker()
worker = Agent(
name="worker",
model=_DirectLlm(model="worker-model"),
model_api_key="test-key",
)
transfer_peer = _workflow_peer(peer_type, tracker)
executed_workflow = ParallelAgent(
name="executed_parallel",
sub_agents=[
SequentialAgent(
name="sequence",
sub_agents=[
_DeterministicAgent(name="sequence_first", tracker=tracker),
_DeterministicAgent(name="sequence_second", tracker=tracker),
],
),
LoopAgent(
name="loop",
sub_agents=[
_DeterministicAgent(
name="loop_stop",
tracker=tracker,
escalate_on_call=1,
)
],
max_iterations=3,
),
],
)
root = Agent(
name="root",
model=_TransferLlm(model="root-model"),
model_api_key="test-key",
sub_agents=[worker, transfer_peer],
)
scenario = SequentialAgent(
name="scenario",
sub_agents=[root, executed_workflow],
)
runner = InMemoryRunner(agent=scenario, app_name="mixed-workflow")
session = await runner.session_service.create_session(
app_name=runner.app_name,
user_id="test-user",
)

events = [
event
async for event in runner.run_async(
user_id="test-user",
session_id=session.id,
new_message=types.UserContent(parts=[types.Part(text="delegate")]),
)
]

assert any(
event.author == "worker"
and event.content
and any(part.text == "worker-ok" for part in (event.content.parts or []))
for event in events
)
assert set(tracker.calls) == {"sequence_first", "sequence_second", "loop_stop"}


@pytest.mark.asyncio
async def test_loop_agent_stops_at_max_iterations_without_escalation() -> None:
tracker = _ExecutionTracker()
Expand Down
30 changes: 29 additions & 1 deletion veadk/utils/patches.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,16 +14,19 @@

import asyncio
import contextvars
from contextlib import asynccontextmanager
import functools
import sys
from contextlib import asynccontextmanager
from typing import Callable

from packaging.version import Version

from veadk.tracing.telemetry.telemetry import (
trace_call_llm,
trace_send_data,
trace_tool_call,
)
from veadk.utils.adk_compat import get_adk_version
from veadk.utils.logger import get_logger
from veadk.version import VERSION

Expand All @@ -35,6 +38,29 @@
)


def patch_adk_workflow_agent_mode() -> None:
"""Backport ADK 2.2's safe peer-transfer check to ADK 2.0 and 2.1.

These releases read ``peer_agent.mode`` unconditionally while collecting
an LLM agent's peer transfer targets. Workflow agents do not define that
field, so mixed LLM/workflow trees fail before the delegated LLM runs. ADK
2.2 fixed the lookup with ``hasattr``; exposing ``None`` on ``BaseAgent``
preserves its target-selection semantics for built-in workflows and
custom agents without replacing ADK internals.
"""
adk_version = get_adk_version()
if not Version("2.0.0") <= adk_version < Version("2.2.0"):
return

try:
from google.adk.agents import BaseAgent

if not hasattr(BaseAgent, "mode"):
BaseAgent.mode = None
except (ImportError, AttributeError) as error:
logger.debug(f"Skip ADK workflow-agent mode patch: {error}")


def _adk_tool_name(tool) -> str:
return (
getattr(tool, "name", None)
Expand Down Expand Up @@ -79,6 +105,8 @@ def patch_asyncio():
- https://github.com/google/adk-python/issues/1429
- https://github.com/google/adk-python/pull/1420
"""
patch_adk_workflow_agent_mode()

original_del = asyncio.base_subprocess.BaseSubprocessTransport.__del__

def patched_del(self):
Expand Down
Loading