From e13f9fba14bbce5f687fe623ea146e8165c57c18 Mon Sep 17 00:00:00 2001 From: Haiyuan Cao Date: Mon, 28 Sep 2026 16:50:59 -0700 Subject: [PATCH 01/29] fix(plugins): record content_formatter failure class in BigQuery rows When BigQueryLoggerConfig.content_formatter raised, or returned a type the parser cannot store, BigQueryAgentAnalyticsPlugin wrote the [FORMATTER_FAILED] sentinel, logged a constant warning, and left the row's error_message NULL, so the developer had no signal about why. Set error_message to a fixed-shape description that names only a class, for example "content_formatter raised ImportError" or "content_formatter returned unsupported type tuple", and never the exception message, args, or traceback, which can embed the content the formatter was protecting. Because type(name, ...) can mint a class named after that content, a name is used only when code chose it: a static type compiled into C, or a class bound under that name in its imported module. Any other class is described by its nearest such ancestor, for example "content_formatter raised ". An event that already carries an error_message, such as a TOOL_ERROR, keeps it first, followed by "; " and the formatter failure. Add BigQueryLoggerConfig.debug_content_formatter_errors (default False), which attaches the traceback to the local formatter-failure warning for debugging. The traceback is never written to BigQuery, and the docstring warns that the process's log handlers can forward it, and the content it embeds, elsewhere. The fail-closed contract is unchanged: the sentinel content, the constant log text by default, and the formatter_failed drop counter. Refs: https://github.com/GoogleCloudPlatform/BigQuery-Agent-Analytics-SDK/issues/485 (item A) Co-Authored-By: Claude Opus 5.5 (1M context) --- .../bigquery_agent_analytics_plugin.py | 128 +++++++- .../test_bigquery_agent_analytics_plugin.py | 295 ++++++++++++++++++ 2 files changed, 415 insertions(+), 8 deletions(-) diff --git a/src/google/adk/plugins/bigquery_agent_analytics_plugin.py b/src/google/adk/plugins/bigquery_agent_analytics_plugin.py index 9024d9543f3..67f972195e5 100644 --- a/src/google/adk/plugins/bigquery_agent_analytics_plugin.py +++ b/src/google/adk/plugins/bigquery_agent_analytics_plugin.py @@ -49,6 +49,7 @@ import random import re +import sys import threading import time from types import MappingProxyType @@ -1012,6 +1013,66 @@ def _sanitize_sensitive_text(text: str, max_len: int) -> tuple[str, bool]: # never fall back to the unformatted payload. _FORMATTER_FAILED_SENTINEL = "[FORMATTER_FAILED]" +# CPython's Py_TPFLAGS_HEAPTYPE: set on every class created at runtime and +# clear on static types compiled into C. +_PY_TPFLAGS_HEAPTYPE = 1 << 9 + + +def _code_defined_class_name(cls: type) -> Optional[str]: + """Returns the qualified name of ``cls`` when code chose it, else None. + + ``type(name, bases, namespace)`` can create a class named after data, such + as the content a formatter was protecting, so a class name is not safe to + record by default. A name counts as chosen by code when ``cls`` is a static + type compiled into C, or when ``cls`` is bound under that name in its + imported module, as a module-level ``class`` statement leaves it and a + runtime ``type()`` call does not. + """ + try: + qualname = str.__str__(cls.__qualname__) + # A metaclass can misreport __flags__, so they are trusted only when the + # metaclass is type itself. + if type(cls) is type and not cls.__flags__ & _PY_TPFLAGS_HEAPTYPE: + return qualname + module = sys.modules.get(cls.__module__) + if module is not None and vars(module).get(qualname) is cls: + return qualname + except Exception: + pass + return None + + +def _formatter_failure_message(cls: type, *, raised: bool) -> str: + """Describes a content_formatter failure for the error_message column. + + Only a class is named: the exception's message, args, and traceback can + embed the content the formatter was protecting. A class whose name was not + chosen by code is described by its nearest ancestor whose name was. + + Args: + cls: The class of the exception the formatter raised, or of the value it + returned when ``raised`` is False. + raised: Whether the formatter raised rather than returned a value. + + Returns: + A fixed-shape message such as ``content_formatter raised ImportError`` or + ``content_formatter returned unsupported type ``. + """ + outcome = "raised" if raised else "returned unsupported type" + name = _code_defined_class_name(cls) + if name is None: + name = "" + try: + for base in cls.__mro__[1:]: + base_name = _code_defined_class_name(base) + if base_name is not None: + name = f"" + break + except Exception: + pass + return f"content_formatter {outcome} {name}" + + # Recursion bound for _recursive_smart_truncate: id()-based cycle detection # cannot catch graphs that create new objects per access (Mock-like duck # typing); the cap turns unbounded recursion into a redacted leaf. @@ -2095,7 +2156,21 @@ class BigQueryLoggerConfig: arriving during the 30-second rotation backoff are also dropped. The unconditional ``event_id`` column remains the deduplication key for default-mode writes. - content_formatter: Optional custom formatter for content. + content_formatter: Optional custom formatter for content, called as + ``content_formatter(content, event_type)``. It is treated as a + redaction boundary, so a failure never falls back to the original + content: if it raises, or returns anything other than a ``str``, + ``dict``, ``list``, ``None``, or an exact ``types.Content``, + ``types.Part``, or ``LlmRequest``, the row is written with content + ``[FORMATTER_FAILED]``, the ``formatter_failed`` counter of + ``get_drop_stats()`` is incremented, and ``error_message`` names the + failure by class only, for example ``content_formatter raised + ImportError``. A class whose name was not chosen by code, such as one + created by ``type(name, ...)``, is named by its nearest ancestor whose + name was, for example ``content_formatter raised ``. An event that already carries an ``error_message``, + such as a ``TOOL_ERROR``, keeps it first, followed by ``; `` and the + formatter failure. gcs_bucket_name: GCS bucket for offloading large content. connection_id: BigQuery connection ID for ObjectRef columns. log_session_metadata: Whether to log session metadata. @@ -2148,6 +2223,18 @@ class BigQueryLoggerConfig: credentials_identifier: Optional explicit string identifier to disambiguate or share background loop states across plugin instances with equivalent credential identities. + debug_content_formatter_errors: When ``True``, an exception raised by + ``content_formatter`` is also logged with its traceback (``exc_info``) + through this module's Python logger, to debug the formatter locally. + The traceback includes the exception message, which can embed the + unformatted content the formatter was protecting, and it reaches every + handler the process has configured: the console, the log file that + ``adk run`` writes, and anything that forwards logs elsewhere, such as + a managed runtime shipping stderr to Cloud Logging. Enable it only + where that content may be seen. The plugin never writes the traceback + to BigQuery; the row's ``error_message`` still names only the + exception class. ``False`` (the default) logs a constant message with + no traceback. """ enabled: bool = True @@ -2233,6 +2320,10 @@ class BigQueryLoggerConfig: # flush_on_run_end is False. use_dedicated_background_loop: Optional[bool] = None credentials_identifier: Optional[str] = None + # Opt-in: attach a failing content_formatter's traceback to the local + # formatter-failure warning. The traceback can embed the unformatted + # content; see the class docstring before enabling it. + debug_content_formatter_errors: bool = False # ============================================================================== @@ -4231,7 +4322,8 @@ def _get_events_schema() -> list[bigquery.SchemaField]: mode="NULLABLE", description=( "Diagnostic message for errors and model termination details;" - " may be populated on LLM_RESPONSE rows whose status is 'OK'." + " may be populated on rows whose status is 'OK', such as" + " LLM_RESPONSE rows and rows whose content_formatter failed." ), ), bigquery.SchemaField( @@ -7183,6 +7275,7 @@ async def _log_event( is_truncated = True timestamp = datetime.now(timezone.utc) + formatter_error: Optional[str] = None if self.config.content_formatter: try: formatted = self.config.content_formatter(raw_content, event_type) @@ -7210,29 +7303,48 @@ async def _log_event( # republish the original content or raise into the safe # callback's traceback log. The # message is CONSTANT: even a class NAME can be payload-derived - # via type(name, ...). + # via type(name, ...), so error_message names only a class whose + # name code chose. logger.warning( "Content formatter returned an unsupported result type for" " event %s; writing sentinel instead of original content.", event_type, ) + formatter_error = _formatter_failure_message( + type(formatted), raised=False + ) formatted = _FORMATTER_FAILED_SENTINEL self._count_local_drop("formatter_failed") raw_content = formatted - except Exception: + except Exception as e: # Fail CLOSED: the formatter is a redaction/privacy # boundary, so its failure must never fall back to the unformatted - # payload. The log message is CONSTANT — the exception message and - # traceback can embed the protected content, and even the class - # NAME can be payload-derived via type(name, ...). + # payload. The log message is CONSTANT and carries no traceback + # unless debug_content_formatter_errors opts in — the exception + # message and traceback can embed the protected content, and even + # the class NAME can be payload-derived via type(name, ...), so + # error_message names only a class whose name code chose. logger.warning( "Content formatter failed for event %s; writing sentinel" " instead of original content.", event_type, + exc_info=self.config.debug_content_formatter_errors, ) + formatter_error = _formatter_failure_message(type(e), raised=True) raw_content = _FORMATTER_FAILED_SENTINEL self._count_local_drop("formatter_failed") + # The event's own diagnostic (e.g. a TOOL_ERROR's exception text) stays + # first and intact so an error row keeps its primary cause; a formatter + # failure is appended after it. + error_message = event_data.error_message + if formatter_error is not None: + error_message = ( + f"{error_message}; {formatter_error}" + if error_message + else formatter_error + ) + trace_id, span_id, parent_span_id = self._resolve_ids( event_data, callback_context ) @@ -7348,7 +7460,7 @@ async def _log_event( "attributes": attributes_json, "latency_ms": latency_json, "status": event_data.status, - "error_message": event_data.error_message, + "error_message": error_message, "is_truncated": is_truncated, } diff --git a/tests/unittests/plugins/test_bigquery_agent_analytics_plugin.py b/tests/unittests/plugins/test_bigquery_agent_analytics_plugin.py index e761e372f98..eb66b94789b 100644 --- a/tests/unittests/plugins/test_bigquery_agent_analytics_plugin.py +++ b/tests/unittests/plugins/test_bigquery_agent_analytics_plugin.py @@ -3813,6 +3813,301 @@ async def test_generation_config_logging( assert attributes.get("labels") == gen_config_kwargs["labels"] +# ============================================================================== +# TEST CLASS: content_formatter failure diagnostics +# ============================================================================== +# Formatters that fail in each way the error_message column must describe. +# They live at module level so that _RedactionServiceError is bound in this +# module the way a library binds its exception classes. + + +class _RedactionServiceError(Exception): + """A module-level exception class, like one a redaction library defines.""" + + +class _ClaimsStaticTypeMeta(type): + """Metaclass whose classes report the type flags of a built-in type.""" + + @property + def __flags__(cls): + return int.__flags__ + + +def _identifier_from_content(content): + """Turns the logged message text into a valid class name.""" + return content.parts[0].text.replace("-", "_") + + +def _raise_import_error(content, event_type): + raise ImportError(f"cannot import name 'redact' (formatting {content})") + + +def _raise_module_level_exception(content, event_type): + raise _RedactionServiceError("redaction backend unavailable") + + +def _raise_local_exception_subclass(content, event_type): + class LocalLookupError(KeyError): + pass + + raise LocalLookupError("missing field") + + +def _raise_payload_named_exception(content, event_type): + raise type(_identifier_from_content(content), (ValueError,), {})() + + +def _raise_payload_named_exception_claiming_static_type(content, event_type): + raise _ClaimsStaticTypeMeta( + _identifier_from_content(content), (ValueError,), {} + )() + + +def _return_tuple(content, event_type): + return ("not", "supported") + + +def _return_generator(content, event_type): + yield content + + +def _return_local_llm_request_subclass(content, event_type): + class LocalRequest(llm_request_lib.LlmRequest): + pass + + return LocalRequest() + + +@pytest.mark.usefixtures( + "mock_auth_default", + "mock_bq_client", + "mock_to_arrow_schema", + "mock_asyncio_to_thread", +) +class TestContentFormatterFailureDiagnostics: + """A failing content_formatter is diagnosable without leaking content. + + The row's error_message names the failure by class. The formatter's input + and the exception's message never reach the row, and the traceback reaches + the local log only when debug_content_formatter_errors is enabled. + """ + + SECRET = "TOPSECRET-4111-1111-1111-1111" + + async def _log_user_message( + self, config, mock_write_client, invocation_context, dummy_arrow_schema + ): + """Logs SECRET as a user message; returns the row and the drop stats.""" + async with managed_plugin( + PROJECT_ID, DATASET_ID, table_id=TABLE_ID, config=config + ) as plugin: + await plugin._ensure_started() + mock_write_client.append_rows.reset_mock() + bigquery_agent_analytics_plugin.TraceManager.push_span(invocation_context) + await plugin.on_user_message_callback( + invocation_context=invocation_context, + user_message=types.Content(parts=[types.Part(text=self.SECRET)]), + ) + await plugin.flush() + row = await _get_captured_event_dict_async( + mock_write_client, dummy_arrow_schema + ) + return row, plugin.get_drop_stats() + + @staticmethod + def _formatter_warnings(caplog): + return [ + record + for record in caplog.records + if record.getMessage().startswith("Content formatter failed") + ] + + @pytest.mark.parametrize( + ("formatter", "expected_error_message"), + [ + pytest.param( + _raise_import_error, + "content_formatter raised ImportError", + id="builtin_exception", + ), + pytest.param( + _raise_module_level_exception, + "content_formatter raised _RedactionServiceError", + id="module_level_exception", + ), + pytest.param( + _raise_local_exception_subclass, + "content_formatter raised ", + id="function_local_exception", + ), + pytest.param( + _raise_payload_named_exception, + "content_formatter raised ", + id="payload_named_exception", + ), + pytest.param( + _raise_payload_named_exception_claiming_static_type, + "content_formatter raised ", + id="payload_named_exception_misreporting_type_flags", + ), + pytest.param( + _return_tuple, + "content_formatter returned unsupported type tuple", + id="unsupported_builtin_result", + ), + pytest.param( + _return_generator, + "content_formatter returned unsupported type generator", + id="unsupported_unexported_builtin_result", + ), + pytest.param( + _return_local_llm_request_subclass, + "content_formatter returned unsupported type" + " ", + id="unsupported_model_subclass_result", + ), + ], + ) + async def test_failed_row_names_the_formatter_failure_by_class( + self, + formatter, + expected_error_message, + mock_write_client, + invocation_context, + dummy_arrow_schema, + ): + """A failed formatter's row fails closed and names the failure's class. + + A class name chosen at runtime, e.g. by type(name, ...) from the content, + is replaced by its nearest code-defined ancestor. + """ + config = bigquery_agent_analytics_plugin.BigQueryLoggerConfig( + content_formatter=formatter + ) + + row, drop_stats = await self._log_user_message( + config, mock_write_client, invocation_context, dummy_arrow_schema + ) + + assert row["error_message"] == expected_error_message + assert ( + row["content"] + == bigquery_agent_analytics_plugin._FORMATTER_FAILED_SENTINEL + ) + assert drop_stats.get("formatter_failed") == 1 + + async def test_formatter_exception_text_never_reaches_the_row( + self, mock_write_client, invocation_context, dummy_arrow_schema + ): + """Neither the exception's message nor the content it embeds is written.""" + config = bigquery_agent_analytics_plugin.BigQueryLoggerConfig( + content_formatter=_raise_import_error + ) + + row, _ = await self._log_user_message( + config, mock_write_client, invocation_context, dummy_arrow_schema + ) + + written = json.dumps(row, default=str) + assert "cannot import name" not in written + assert self.SECRET not in written + + @pytest.mark.parametrize( + ("tool_error_text", "expected_error_message"), + [ + pytest.param( + "upstream timed out after 30s", + "upstream timed out after 30s;" + " content_formatter raised ImportError", + id="appended_after_existing_message", + ), + pytest.param( + "", + "content_formatter raised ImportError", + id="empty_existing_message", + ), + ], + ) + async def test_formatter_failure_follows_the_events_own_error_message( + self, + tool_error_text, + expected_error_message, + mock_write_client, + tool_context, + dummy_arrow_schema, + ): + """An error row keeps its own diagnostic first; the formatter's follows.""" + config = bigquery_agent_analytics_plugin.BigQueryLoggerConfig( + content_formatter=_raise_import_error + ) + tool = mock.create_autospec( + base_tool_lib.BaseTool, instance=True, spec_set=True + ) + type(tool).name = mock.PropertyMock(return_value="lookup") + + async with managed_plugin( + PROJECT_ID, DATASET_ID, table_id=TABLE_ID, config=config + ) as plugin: + await plugin._ensure_started() + mock_write_client.append_rows.reset_mock() + bigquery_agent_analytics_plugin.TraceManager.push_span(tool_context) + await plugin.on_tool_error_callback( + tool=tool, + tool_args={"account": self.SECRET}, + tool_context=tool_context, + error=RuntimeError(tool_error_text), + ) + await plugin.flush() + row = await _get_captured_event_dict_async( + mock_write_client, dummy_arrow_schema + ) + + assert row["error_message"] == expected_error_message + + async def test_formatter_traceback_is_not_logged_by_default( + self, mock_write_client, invocation_context, dummy_arrow_schema, caplog + ): + """By default the formatter-failure warning carries no traceback.""" + config = bigquery_agent_analytics_plugin.BigQueryLoggerConfig( + content_formatter=_raise_import_error + ) + + with caplog.at_level(logging.WARNING): + await self._log_user_message( + config, mock_write_client, invocation_context, dummy_arrow_schema + ) + + warnings = self._formatter_warnings(caplog) + assert len(warnings) == 1 + assert not warnings[0].exc_info + assert self.SECRET not in caplog.text + + async def test_debug_flag_logs_traceback_locally_but_not_to_the_row( + self, mock_write_client, invocation_context, dummy_arrow_schema, caplog + ): + """debug_content_formatter_errors sends the traceback to the log only. + + The traceback carries the exception message and the content it embeds, + so the row still names only the exception class. + """ + config = bigquery_agent_analytics_plugin.BigQueryLoggerConfig( + content_formatter=_raise_import_error, + debug_content_formatter_errors=True, + ) + + with caplog.at_level(logging.WARNING): + row, _ = await self._log_user_message( + config, mock_write_client, invocation_context, dummy_arrow_schema + ) + + warnings = self._formatter_warnings(caplog) + assert len(warnings) == 1 + assert warnings[0].exc_info[0] is ImportError + assert self.SECRET in caplog.text + assert row["error_message"] == "content_formatter raised ImportError" + assert self.SECRET not in json.dumps(row, default=str) + + class TestSafeCallbackDecorator: """Tests that _safe_callback prevents plugin errors from propagating.""" From def458b609c2811d137b0332b2fc7b201dddd5c0 Mon Sep 17 00:00:00 2001 From: feiiiiii5 Date: Mon, 28 Sep 2026 21:00:41 -0700 Subject: [PATCH 02/29] fix: stamp Redis sessions with the event timestamp Merge https://github.com/google/adk-python/pull/7293 Fixes #7292 PiperOrigin-RevId: 990023711 --- .../redis/_redis_session_service.py | 7 +++- .../redis/test_redis_session_service.py | 42 +++++++++++++++++++ tests/unittests/sessions/_conformance.py | 4 -- 3 files changed, 48 insertions(+), 5 deletions(-) diff --git a/src/google/adk/integrations/redis/_redis_session_service.py b/src/google/adk/integrations/redis/_redis_session_service.py index cc7c41c7cd8..db493a1e6ca 100644 --- a/src/google/adk/integrations/redis/_redis_session_service.py +++ b/src/google/adk/integrations/redis/_redis_session_service.py @@ -340,7 +340,12 @@ async def append_event(self, session: Session, event: Event) -> Event: """Appends an event to the session and synchronizes state in Redis.""" client = self._get_redis() event = await super().append_event(session, event) - session.last_update_time = time.time() + # Stamp the session with the event's own timestamp, matching + # InMemorySessionService, SqliteSessionService and DatabaseSessionService. + # The wall clock is wrong here: an event records when it was produced, which + # can predate the append when it is replayed or re-delivered, and + # last_update_time is the key list_sessions orders by. + session.last_update_time = event.timestamp # Sync app and user state deltas to their respective keys if event.actions and event.actions.state_delta: diff --git a/tests/unittests/integrations/redis/test_redis_session_service.py b/tests/unittests/integrations/redis/test_redis_session_service.py index 72d8438be38..e617403049f 100644 --- a/tests/unittests/integrations/redis/test_redis_session_service.py +++ b/tests/unittests/integrations/redis/test_redis_session_service.py @@ -380,6 +380,48 @@ async def test_append_event_and_state_delta(session_service): assert fetched.state["app:status"] == "active" +@pytest.mark.asyncio +async def test_append_event_stamps_session_with_event_timestamp( + session_service, +): + """The session records when the event happened, not when it was appended. + + `last_update_time` is what `list_sessions` orders by, so stamping it with the + wall clock makes an event that is replayed, re-delivered or imported push a + session forward to its append time instead of its own. Every other backend + stores `event.timestamp`; the shared contract test asserts the same. + """ + session = await session_service.create_session( + app_name="app1", + user_id="u1", + ) + + event_timestamp = session.last_update_time + 10 + event = Event( + author="agent", + invocation_id="inv1", + timestamp=event_timestamp, + ) + + # Pin the wall clock far from the event's own timestamp so the current + # implementation cannot agree with the expected value by coincidence. + with mock.patch( + "google.adk.integrations.redis._redis_session_service.time" + ) as clock: + clock.time.return_value = event_timestamp + 100 + await session_service.append_event(session, event) + + assert session.last_update_time == pytest.approx(event_timestamp, abs=1e-6) + + fetched = await session_service.get_session( + app_name="app1", + user_id="u1", + session_id=session.id, + ) + assert fetched is not None + assert fetched.last_update_time == pytest.approx(event_timestamp, abs=1e-6) + + @pytest.mark.asyncio async def test_app_and_user_state_ttl(fake_redis, session_service): await session_service.create_session( diff --git a/tests/unittests/sessions/_conformance.py b/tests/unittests/sessions/_conformance.py index 3a62bec21da..842120f611e 100644 --- a/tests/unittests/sessions/_conformance.py +++ b/tests/unittests/sessions/_conformance.py @@ -144,10 +144,6 @@ async def _make_per_agent_database( 'redis', _make_redis, divergences={ - 'test_session_last_update_time_updates_on_event': ( - 'Redis stamps the session with the wall clock instead of the' - " appended event's timestamp." - ), 'test_append_event_to_unknown_session_raises_session_not_found': ( 'Redis writes the session key unconditionally on append, so' ' appending to a session it has never stored creates one' From 2b9fbd0f8883404a35acfd61e18144340bd480a8 Mon Sep 17 00:00:00 2001 From: Haiyuan Cao Date: Mon, 28 Sep 2026 21:05:29 -0700 Subject: [PATCH 03/29] fix(plugins): stop formatter diagnostics leaking names or dropping rows Review of the previous commit found three ways the new content_formatter diagnostics could still leak payload text or drop a row. This closes them. Trusted labels only. A class created at runtime can be named after the content a formatter was protecting and then bound into its module, so "the class is bound under its name in its module" proved nothing, and such a name reached error_message. error_message now names a class only by a label that no runtime data can have chosen: the compiled name of a static C type, or a fixed string for an allowlisted class matched by identity (LlmRequest, types.Content, types.Part, pydantic BaseModel, and google.api_core GoogleAPICallError). Every other class, including every Python-defined library exception and every class the developer defines, is described by its nearest trusted ancestor, for example "content_formatter raised ". Trade-off: fewer exact names, for example json.JSONDecodeError now reads "" and a google.api_core NotFound ""; debug_content_formatter_errors shows the exact class locally. No class code runs while diagnosing. Flags, name, and MRO are read through type's own descriptors, and allowlist matching compares identity, so a metaclass __getattribute__, __eq__, or __hash__ can no longer run inside the fail-closed boundary. Such a hook could raise asyncio.CancelledError, which escaped the boundary's `except Exception` and dropped the row. Classification now cannot raise at all, so it needs no BaseException handler, and a genuine KeyboardInterrupt or SystemExit delivered by a signal still propagates. The exact-class check on formatter results also compares identity now, for the same reason. Best-effort debug traceback. With debug_content_formatter_errors, the traceback is rendered to text once, inside the plugin, and appended to the warning; handlers never receive the live exception, whose own code could otherwise fail inside a stock StreamHandler and drop the row. A rendering failure falls back to a constant placeholder. KeyboardInterrupt and SystemExit still propagate, because a signal can deliver them at any bytecode; any other BaseException raised while rendering, CancelledError included, can only come from the exception being rendered, since rendering never awaits. The warning is emitted after the except block, so a failing handler's handleError cannot reach the formatter's exception through its exception chain. Tests pin each guard: registered payload-named classes and their subclasses, a renamed allowlisted class, metaclass hooks raising CancelledError, SystemExit, or RuntimeError on both failure paths, unrenderable tracebacks under a stock StreamHandler, interrupt propagation, and a failing log handler. Removing any guard fails at least one test. Refs: https://github.com/GoogleCloudPlatform/BigQuery-Agent-Analytics-SDK/issues/485 (item A) Co-Authored-By: Claude Opus 5.5 (1M context) --- .../bigquery_agent_analytics_plugin.py | 216 +++++--- .../test_bigquery_agent_analytics_plugin.py | 464 +++++++++++++++++- 2 files changed, 591 insertions(+), 89 deletions(-) diff --git a/src/google/adk/plugins/bigquery_agent_analytics_plugin.py b/src/google/adk/plugins/bigquery_agent_analytics_plugin.py index 67f972195e5..97938f93105 100644 --- a/src/google/adk/plugins/bigquery_agent_analytics_plugin.py +++ b/src/google/adk/plugins/bigquery_agent_analytics_plugin.py @@ -49,7 +49,6 @@ import random import re -import sys import threading import time from types import MappingProxyType @@ -87,6 +86,7 @@ import google.cloud.storage as cloud_storage from google.genai import types from opentelemetry import trace +from pydantic import BaseModel try: import pyarrow as pa @@ -1014,40 +1014,62 @@ def _sanitize_sensitive_text(text: str, max_len: int) -> tuple[str, bool]: _FORMATTER_FAILED_SENTINEL = "[FORMATTER_FAILED]" # CPython's Py_TPFLAGS_HEAPTYPE: set on every class created at runtime and -# clear on static types compiled into C. +# clear on static types compiled into C, whose names cannot be reassigned. _PY_TPFLAGS_HEAPTYPE = 1 << 9 +# type's own descriptors. Calling them directly reads a class's flags, name, +# and MRO without running code the class controls: an ordinary attribute read +# goes through the metaclass, whose hooks can lie or raise anything, including +# BaseException subclasses that the plugin's boundaries deliberately let pass. +_TYPE_FLAGS = type.__dict__["__flags__"] +_TYPE_NAME = type.__dict__["__name__"] +_TYPE_MRO = type.__dict__["__mro__"] + +# Runtime-created classes that a formatter commonly raises or returns, each +# with a fixed label. Labels are never read from the class, because a class +# created at runtime can be renamed. +_TRUSTED_CLASS_LABELS: tuple[tuple[type, str], ...] = ( + (LlmRequest, "LlmRequest"), + (types.Content, "Content"), + (types.Part, "Part"), + (BaseModel, "BaseModel"), + (api_exceptions.GoogleAPICallError, "GoogleAPICallError"), +) + + +def _trusted_class_label(cls: type) -> Optional[str]: + """Returns a label for ``cls`` that no runtime data can have chosen. + + Only static types compiled into C, whose names are fixed when the + interpreter or extension is built, and the classes in + ``_TRUSTED_CLASS_LABELS`` have one. Any other class can be created at + runtime by ``type(name, bases, namespace)`` with a name taken from the + content a formatter was protecting, bound into a module under that name, or + renamed, so its name is never trusted, wherever it is defined. -def _code_defined_class_name(cls: type) -> Optional[str]: - """Returns the qualified name of ``cls`` when code chose it, else None. + Args: + cls: The class to label. - ``type(name, bases, namespace)`` can create a class named after data, such - as the content a formatter was protecting, so a class name is not safe to - record by default. A name counts as chosen by code when ``cls`` is a static - type compiled into C, or when ``cls`` is bound under that name in its - imported module, as a module-level ``class`` statement leaves it and a - runtime ``type()`` call does not. + Returns: + The label, or None when ``cls`` has none. """ - try: - qualname = str.__str__(cls.__qualname__) - # A metaclass can misreport __flags__, so they are trusted only when the - # metaclass is type itself. - if type(cls) is type and not cls.__flags__ & _PY_TPFLAGS_HEAPTYPE: - return qualname - module = sys.modules.get(cls.__module__) - if module is not None and vars(module).get(qualname) is cls: - return qualname - except Exception: - pass + for trusted, label in _TRUSTED_CLASS_LABELS: + if cls is trusted: + return label + if not _TYPE_FLAGS.__get__(cls) & _PY_TPFLAGS_HEAPTYPE: + name: str = _TYPE_NAME.__get__(cls) + return name return None def _formatter_failure_message(cls: type, *, raised: bool) -> str: """Describes a content_formatter failure for the error_message column. - Only a class is named: the exception's message, args, and traceback can - embed the content the formatter was protecting. A class whose name was not - chosen by code is described by its nearest ancestor whose name was. + The class is named only by a trusted label, and a class without one by its + nearest ancestor that has one. The exception's message, args, and traceback + are never used: they can embed the content the formatter was protecting. + Nothing here runs code the class controls, so describing a failure cannot + raise, whatever the failure's class does. Args: cls: The class of the exception the formatter raised, or of the value it @@ -1059,18 +1081,41 @@ def _formatter_failure_message(cls: type, *, raised: bool) -> str: ``content_formatter returned unsupported type ``. """ outcome = "raised" if raised else "returned unsupported type" - name = _code_defined_class_name(cls) - if name is None: - name = "" - try: - for base in cls.__mro__[1:]: - base_name = _code_defined_class_name(base) - if base_name is not None: - name = f"" - break - except Exception: - pass - return f"content_formatter {outcome} {name}" + for depth, ancestor in enumerate(_TYPE_MRO.__get__(cls)): + label = _trusted_class_label(ancestor) + if label is not None: + if depth == 0: + return f"content_formatter {outcome} {label}" + return f"content_formatter {outcome} " + # Unreachable for an instantiable class, whose MRO ends with object. + return f"content_formatter {outcome} " + + +def _render_formatter_traceback(error: BaseException) -> str: + """Renders a content_formatter exception's traceback for debug logging. + + Rendering runs the exception's own code: its ``__str__`` and the attribute + hooks that expose its traceback and chained exceptions. Whatever that code + raises is contained here and replaced by a constant, so rendering can never + affect the row, with two exceptions. KeyboardInterrupt and SystemExit + propagate, as they do everywhere else in the plugin, because a signal + handler can deliver either at any bytecode, so one raised here may be + genuine. Any other BaseException, asyncio.CancelledError included, can only + come from the exception being rendered: rendering never awaits, so it cannot + receive a real cancellation. + + Args: + error: The exception the formatter raised. + + Returns: + The rendered traceback, or ``[traceback could not be rendered]``. + """ + try: + return "".join(traceback_module.format_exception(error)).rstrip("\n") + except (KeyboardInterrupt, SystemExit): + raise + except BaseException: + return "[traceback could not be rendered]" # Recursion bound for _recursive_smart_truncate: id()-based cycle detection @@ -2165,12 +2210,15 @@ class BigQueryLoggerConfig: ``[FORMATTER_FAILED]``, the ``formatter_failed`` counter of ``get_drop_stats()`` is incremented, and ``error_message`` names the failure by class only, for example ``content_formatter raised - ImportError``. A class whose name was not chosen by code, such as one - created by ``type(name, ...)``, is named by its nearest ancestor whose - name was, for example ``content_formatter raised ``. An event that already carries an ``error_message``, - such as a ``TOOL_ERROR``, keeps it first, followed by ``; `` and the - formatter failure. + ImportError``. Because a class can be created or renamed at runtime + with a name taken from the content, only built-in types and a few + trusted classes (``LlmRequest``, ``types.Content``, ``types.Part``, + pydantic ``BaseModel``, and ``google.api_core`` ``GoogleAPICallError``) + are named. Any other class, including one your own code defines, is + described by its nearest named ancestor, for example + ``content_formatter raised ``. An event that + already carries an ``error_message``, such as a ``TOOL_ERROR``, keeps + it first, followed by ``; `` and the formatter failure. gcs_bucket_name: GCS bucket for offloading large content. connection_id: BigQuery connection ID for ObjectRef columns. log_session_metadata: Whether to log session metadata. @@ -2223,18 +2271,21 @@ class BigQueryLoggerConfig: credentials_identifier: Optional explicit string identifier to disambiguate or share background loop states across plugin instances with equivalent credential identities. - debug_content_formatter_errors: When ``True``, an exception raised by - ``content_formatter`` is also logged with its traceback (``exc_info``) - through this module's Python logger, to debug the formatter locally. - The traceback includes the exception message, which can embed the - unformatted content the formatter was protecting, and it reaches every - handler the process has configured: the console, the log file that - ``adk run`` writes, and anything that forwards logs elsewhere, such as - a managed runtime shipping stderr to Cloud Logging. Enable it only - where that content may be seen. The plugin never writes the traceback - to BigQuery; the row's ``error_message`` still names only the - exception class. ``False`` (the default) logs a constant message with - no traceback. + debug_content_formatter_errors: When ``True``, the traceback of an + exception raised by ``content_formatter`` is rendered to text and + appended to the formatter-failure warning that this module's Python + logger emits, to debug the formatter locally. The traceback includes + the exception message, which can embed the unformatted content the + formatter was protecting, and it reaches every handler the process + has configured: the console, the log file that ``adk run`` writes, and + anything that forwards logs elsewhere, such as a managed runtime + shipping stderr to Cloud Logging. Enable it only where that content + may be seen. Rendering is best effort: if the exception's own code + fails while it is rendered, a constant placeholder is logged instead, + and the row is unaffected. The plugin never writes the traceback to + BigQuery; the row's ``error_message`` still names only the exception + class. ``False`` (the default) logs a constant message with no + traceback. """ enabled: bool = True @@ -2320,8 +2371,8 @@ class BigQueryLoggerConfig: # flush_on_run_end is False. use_dedicated_background_loop: Optional[bool] = None credentials_identifier: Optional[str] = None - # Opt-in: attach a failing content_formatter's traceback to the local - # formatter-failure warning. The traceback can embed the unformatted + # Opt-in: append a failing content_formatter's rendered traceback to the + # local formatter-failure warning. The traceback can embed the unformatted # content; see the class docstring before enabling it. debug_content_formatter_errors: bool = False @@ -7276,6 +7327,8 @@ async def _log_event( timestamp = datetime.now(timezone.utc) formatter_error: Optional[str] = None + formatter_raised = False + formatter_traceback: Optional[str] = None if self.config.content_formatter: try: formatted = self.config.content_formatter(raw_content, event_type) @@ -7294,7 +7347,12 @@ async def _log_event( # safe callback's traceback log. dict/list subclasses stay isinstance-based — the # parser routes them through the hardened recursive # sanitizer, whose protocol boundary already fails closed. - type(formatted) in (types.Content, types.Part, LlmRequest) + # Compared by identity: `in` would call the result class's + # metaclass __eq__, which can raise anything or claim a match. + any( + type(formatted) is shape + for shape in (types.Content, types.Part, LlmRequest) + ) or isinstance(formatted, (dict, list)) ): # The formatter is typed Any: a non-native result would reach @@ -7303,8 +7361,8 @@ async def _log_event( # republish the original content or raise into the safe # callback's traceback log. The # message is CONSTANT: even a class NAME can be payload-derived - # via type(name, ...), so error_message names only a class whose - # name code chose. + # via type(name, ...), so error_message uses only a trusted class + # label. logger.warning( "Content formatter returned an unsupported result type for" " event %s; writing sentinel instead of original content.", @@ -7319,24 +7377,40 @@ async def _log_event( except Exception as e: # Fail CLOSED: the formatter is a redaction/privacy # boundary, so its failure must never fall back to the unformatted - # payload. The log message is CONSTANT and carries no traceback - # unless debug_content_formatter_errors opts in — the exception - # message and traceback can embed the protected content, and even - # the class NAME can be payload-derived via type(name, ...), so - # error_message names only a class whose name code chose. - logger.warning( - "Content formatter failed for event %s; writing sentinel" - " instead of original content.", - event_type, - exc_info=self.config.debug_content_formatter_errors, - ) + # payload. The exception message and traceback can embed the + # protected content, and even the class NAME can be payload-derived + # via type(name, ...), so error_message uses only a trusted class + # label and the warning below stays CONSTANT. The traceback is + # rendered, best effort, only when debug_content_formatter_errors + # opts in; no handler ever receives the live exception. + formatter_raised = True formatter_error = _formatter_failure_message(type(e), raised=True) + if self.config.debug_content_formatter_errors: + formatter_traceback = _render_formatter_traceback(e) raw_content = _FORMATTER_FAILED_SENTINEL self._count_local_drop("formatter_failed") + if formatter_raised: + # Logged only once the formatter's exception is no longer being + # handled: a failing handler's handleError prints its own exception + # chain, which would otherwise reach that exception and run its code. + if formatter_traceback is None: + logger.warning( + "Content formatter failed for event %s; writing sentinel" + " instead of original content.", + event_type, + ) + else: + logger.warning( + "Content formatter failed for event %s; writing sentinel" + " instead of original content. Debug traceback:\n%s", + event_type, + formatter_traceback, + ) # The event's own diagnostic (e.g. a TOOL_ERROR's exception text) stays # first and intact so an error row keeps its primary cause; a formatter - # failure is appended after it. + # failure is appended after it. The note skips the bounded sanitizer + # above: it is fixed text and a trusted class label, never free text. error_message = event_data.error_message if formatter_error is not None: error_message = ( diff --git a/tests/unittests/plugins/test_bigquery_agent_analytics_plugin.py b/tests/unittests/plugins/test_bigquery_agent_analytics_plugin.py index eb66b94789b..9ba6e882286 100644 --- a/tests/unittests/plugins/test_bigquery_agent_analytics_plugin.py +++ b/tests/unittests/plugins/test_bigquery_agent_analytics_plugin.py @@ -17,6 +17,7 @@ import concurrent.futures import contextlib import dataclasses +import io import json import logging import os @@ -3817,8 +3818,8 @@ async def test_generation_config_logging( # TEST CLASS: content_formatter failure diagnostics # ============================================================================== # Formatters that fail in each way the error_message column must describe. -# They live at module level so that _RedactionServiceError is bound in this -# module the way a library binds its exception classes. +# Payload-derived class names are built from the logged message text, so a +# leak shows up as that text in a column or a log line. class _RedactionServiceError(Exception): @@ -3833,11 +3834,91 @@ def __flags__(cls): return int.__flags__ +class _NeitherInterruptNorCancellation(BaseException): + """A BaseException that is not KeyboardInterrupt, SystemExit, or cancel.""" + + def _identifier_from_content(content): """Turns the logged message text into a valid class name.""" return content.parts[0].text.replace("-", "_") +def _register_payload_named_class(content, bases): + """Creates a payload-named class and binds it in this module under its name. + + Class factories do this so that pickling can find their products, which is + what a name-in-its-module check would take as a class defined by code. + """ + name = _identifier_from_content(content) + cls = type(name, bases, {"__module__": __name__}) + globals()[name] = cls + return cls + + +class _Tripwire: + """Makes hostile hooks raise only while armed. + + pytest reads an escaped exception's class name and message when it reports + a failure. Hooks that still raised then would abort the whole session, so + each test disarms its tripwire before pytest reports anything. + """ + + def __init__(self, error_type): + self.error_type = error_type + self.armed = False + + def fire(self, hook): + if self.armed: + raise self.error_type(f"TRIPWIRE: {hook} ran") + + +def _metaclass_whose_hooks_raise(tripwire): + """Returns a metaclass whose attribute, equality, and hash hooks fire.""" + + class _HookedMeta(type): + + def __getattribute__(cls, name): + tripwire.fire(f"metaclass __getattribute__({name!r})") + return super().__getattribute__(name) + + def __eq__(cls, other): + tripwire.fire("metaclass __eq__") + return super().__eq__(other) + + def __hash__(cls): + tripwire.fire("metaclass __hash__") + return super().__hash__() + + return _HookedMeta + + +def _unrenderable_exception(tripwire): + """Returns an exception whose traceback cannot be rendered while armed. + + Rendering reads the traceback and the chained exceptions through the + exception's own __getattribute__ and calls its __str__; both fire here. + """ + + class _UnrenderableError(ValueError): + + def __getattribute__(self, name): + if name in ( + "__traceback__", + "__cause__", + "__context__", + "__suppress_context__", + "__notes__", + ): + tripwire.fire(f"exception __getattribute__({name!r})") + return super().__getattribute__(name) + + def __str__(self): + tripwire.fire("exception __str__") + return "unrenderable" + + return _UnrenderableError() + + def _raise_import_error(content, event_type): raise ImportError(f"cannot import name 'redact' (formatting {content})") @@ -3863,6 +3944,19 @@ def _raise_payload_named_exception_claiming_static_type(content, event_type): )() +def _raise_registered_payload_named_exception(content, event_type): + raise _register_payload_named_class(content, (ValueError,))() + + +def _raise_subclass_of_registered_payload_named_exception(content, event_type): + registered = _register_payload_named_class(content, (ValueError,)) + raise type("Unregistered", (registered,), {})() + + +def _raise_google_api_error(content, event_type): + raise api_exceptions.NotFound("redaction template not found") + + def _return_tuple(content, event_type): return ("not", "supported") @@ -3871,6 +3965,10 @@ def _return_generator(content, event_type): yield content +def _return_registered_payload_named_object(content, event_type): + return _register_payload_named_class(content, ())() + + def _return_local_llm_request_subclass(content, event_type): class LocalRequest(llm_request_lib.LlmRequest): pass @@ -3878,6 +3976,27 @@ class LocalRequest(llm_request_lib.LlmRequest): return LocalRequest() +def _return_local_content_subclass(content, event_type): + class LocalContent(types.Content): + pass + + return LocalContent() + + +def _return_local_part_subclass(content, event_type): + class LocalPart(types.Part): + pass + + return LocalPart() + + +def _return_pydantic_model(content, event_type): + class RedactedPayload(BaseModel): + text: str = "[REDACTED]" + + return RedactedPayload() + + @pytest.mark.usefixtures( "mock_auth_default", "mock_bq_client", @@ -3887,12 +4006,24 @@ class LocalRequest(llm_request_lib.LlmRequest): class TestContentFormatterFailureDiagnostics: """A failing content_formatter is diagnosable without leaking content. - The row's error_message names the failure by class. The formatter's input - and the exception's message never reach the row, and the traceback reaches - the local log only when debug_content_formatter_errors is enabled. + The row's error_message names the failure by a trusted class label. The + formatter's input, the exception's message, and any class name taken from + the class itself never reach the row. Diagnosing the failure never drops + the row, and the traceback reaches the local log only when + debug_content_formatter_errors is enabled. """ SECRET = "TOPSECRET-4111-1111-1111-1111" + PAYLOAD_IDENTIFIER = "TOPSECRET_4111_1111_1111_1111" + DEFAULT_WARNING = ( + "Content formatter failed for event USER_MESSAGE_RECEIVED; writing" + " sentinel instead of original content." + ) + + @pytest.fixture(autouse=True) + def _unbind_registered_payload_classes(self): + yield + globals().pop(self.PAYLOAD_IDENTIFIER, None) async def _log_user_message( self, config, mock_write_client, invocation_context, dummy_arrow_schema @@ -3922,6 +4053,23 @@ def _formatter_warnings(caplog): if record.getMessage().startswith("Content formatter failed") ] + @staticmethod + @contextlib.contextmanager + def _standard_handler_on_plugin_logger(stream): + """Attaches a stock logging.StreamHandler, as an application would.""" + plugin_logger = logging.getLogger( + "google_adk." + bigquery_agent_analytics_plugin.__name__ + ) + handler = logging.StreamHandler(stream) + previous_level = plugin_logger.level + plugin_logger.addHandler(handler) + plugin_logger.setLevel(logging.WARNING) + try: + yield + finally: + plugin_logger.removeHandler(handler) + plugin_logger.setLevel(previous_level) + @pytest.mark.parametrize( ("formatter", "expected_error_message"), [ @@ -3932,7 +4080,7 @@ def _formatter_warnings(caplog): ), pytest.param( _raise_module_level_exception, - "content_formatter raised _RedactionServiceError", + "content_formatter raised ", id="module_level_exception", ), pytest.param( @@ -3950,6 +4098,21 @@ def _formatter_warnings(caplog): "content_formatter raised ", id="payload_named_exception_misreporting_type_flags", ), + pytest.param( + _raise_registered_payload_named_exception, + "content_formatter raised ", + id="registered_payload_named_exception", + ), + pytest.param( + _raise_subclass_of_registered_payload_named_exception, + "content_formatter raised ", + id="subclass_of_registered_payload_named_exception", + ), + pytest.param( + _raise_google_api_error, + "content_formatter raised ", + id="google_api_error", + ), pytest.param( _return_tuple, "content_formatter returned unsupported type tuple", @@ -3960,11 +4123,34 @@ def _formatter_warnings(caplog): "content_formatter returned unsupported type generator", id="unsupported_unexported_builtin_result", ), + pytest.param( + _return_registered_payload_named_object, + "content_formatter returned unsupported type" + " ", + id="registered_payload_named_result", + ), pytest.param( _return_local_llm_request_subclass, "content_formatter returned unsupported type" " ", - id="unsupported_model_subclass_result", + id="llm_request_subclass_result", + ), + pytest.param( + _return_local_content_subclass, + "content_formatter returned unsupported type" + " ", + id="content_subclass_result", + ), + pytest.param( + _return_local_part_subclass, + "content_formatter returned unsupported type ", + id="part_subclass_result", + ), + pytest.param( + _return_pydantic_model, + "content_formatter returned unsupported type" + " ", + id="pydantic_model_result", ), ], ) @@ -3975,19 +4161,116 @@ async def test_failed_row_names_the_formatter_failure_by_class( mock_write_client, invocation_context, dummy_arrow_schema, + caplog, ): - """A failed formatter's row fails closed and names the failure's class. + """A failed formatter's row fails closed and names a trusted class. - A class name chosen at runtime, e.g. by type(name, ...) from the content, - is replaced by its nearest code-defined ancestor. + Only a built-in type or an allowlisted class is named. Any other class, + including a module-level one or one named after the content, is + described by its nearest such ancestor, and its own name appears in no + column and no default log line. """ config = bigquery_agent_analytics_plugin.BigQueryLoggerConfig( content_formatter=formatter ) - row, drop_stats = await self._log_user_message( - config, mock_write_client, invocation_context, dummy_arrow_schema + with caplog.at_level(logging.WARNING): + row, drop_stats = await self._log_user_message( + config, mock_write_client, invocation_context, dummy_arrow_schema + ) + + assert row["error_message"] == expected_error_message + assert ( + row["content"] + == bigquery_agent_analytics_plugin._FORMATTER_FAILED_SENTINEL ) + assert drop_stats.get("formatter_failed") == 1 + written = json.dumps(row, default=str) + for payload_text in (self.SECRET, self.PAYLOAD_IDENTIFIER): + assert payload_text not in written + assert payload_text not in caplog.text + + async def test_trusted_class_label_is_fixed_text_not_its_current_name( + self, mock_write_client, invocation_context, dummy_arrow_schema + ): + """Renaming an allowlisted class at runtime cannot change its label.""" + trusted = api_exceptions.GoogleAPICallError + original_names = (trusted.__name__, trusted.__qualname__) + config = bigquery_agent_analytics_plugin.BigQueryLoggerConfig( + content_formatter=_raise_google_api_error + ) + + trusted.__name__ = trusted.__qualname__ = self.PAYLOAD_IDENTIFIER + try: + row, _ = await self._log_user_message( + config, mock_write_client, invocation_context, dummy_arrow_schema + ) + finally: + trusted.__name__, trusted.__qualname__ = original_names + + assert row["error_message"] == ( + "content_formatter raised " + ) + + @pytest.mark.parametrize( + "hook_error", + [asyncio.CancelledError, SystemExit, RuntimeError], + ids=["cancelled_error", "system_exit", "runtime_error"], + ) + @pytest.mark.parametrize( + ("raised", "expected_error_message"), + [ + pytest.param( + True, + "content_formatter raised ", + id="raised", + ), + pytest.param( + False, + "content_formatter returned unsupported type" + " ", + id="returned", + ), + ], + ) + async def test_naming_the_failure_runs_none_of_the_class_hooks( + self, + raised, + expected_error_message, + hook_error, + mock_write_client, + invocation_context, + dummy_arrow_schema, + caplog, + ): + """Diagnosing a failure never runs the failed class's metaclass hooks. + + Those hooks can raise anything, including BaseException subclasses that + the fail-closed boundary deliberately lets through, so the row, its + sentinel, and the drop counter must not depend on them. + """ + tripwire = _Tripwire(hook_error) + hooked_meta = _metaclass_whose_hooks_raise(tripwire) + + def formatter(content, event_type): + name = _identifier_from_content(content) + failure = hooked_meta(name, (ValueError,) if raised else (), {})() + tripwire.armed = True + if raised: + raise failure + return failure + + config = bigquery_agent_analytics_plugin.BigQueryLoggerConfig( + content_formatter=formatter + ) + + try: + with caplog.at_level(logging.WARNING): + row, drop_stats = await self._log_user_message( + config, mock_write_client, invocation_context, dummy_arrow_schema + ) + finally: + tripwire.armed = False assert row["error_message"] == expected_error_message assert ( @@ -3995,6 +4278,7 @@ async def test_failed_row_names_the_formatter_failure_by_class( == bigquery_agent_analytics_plugin._FORMATTER_FAILED_SENTINEL ) assert drop_stats.get("formatter_failed") == 1 + assert "TRIPWIRE" not in caplog.text async def test_formatter_exception_text_never_reaches_the_row( self, mock_write_client, invocation_context, dummy_arrow_schema @@ -4067,7 +4351,7 @@ async def test_formatter_failure_follows_the_events_own_error_message( async def test_formatter_traceback_is_not_logged_by_default( self, mock_write_client, invocation_context, dummy_arrow_schema, caplog ): - """By default the formatter-failure warning carries no traceback.""" + """By default the formatter-failure warning is constant, no traceback.""" config = bigquery_agent_analytics_plugin.BigQueryLoggerConfig( content_formatter=_raise_import_error ) @@ -4078,7 +4362,9 @@ async def test_formatter_traceback_is_not_logged_by_default( ) warnings = self._formatter_warnings(caplog) - assert len(warnings) == 1 + assert [record.getMessage() for record in warnings] == [ + self.DEFAULT_WARNING + ] assert not warnings[0].exc_info assert self.SECRET not in caplog.text @@ -4087,8 +4373,9 @@ async def test_debug_flag_logs_traceback_locally_but_not_to_the_row( ): """debug_content_formatter_errors sends the traceback to the log only. - The traceback carries the exception message and the content it embeds, - so the row still names only the exception class. + The traceback is rendered to text before logging, so no handler ever + receives the live exception. It carries the exception message and the + content that message embeds, so the row still names only the class. """ config = bigquery_agent_analytics_plugin.BigQueryLoggerConfig( content_formatter=_raise_import_error, @@ -4102,11 +4389,152 @@ async def test_debug_flag_logs_traceback_locally_but_not_to_the_row( warnings = self._formatter_warnings(caplog) assert len(warnings) == 1 - assert warnings[0].exc_info[0] is ImportError - assert self.SECRET in caplog.text + message = warnings[0].getMessage() + assert message.startswith(self.DEFAULT_WARNING) + assert "Traceback (most recent call last)" in message + assert self.SECRET in message + assert not warnings[0].exc_info assert row["error_message"] == "content_formatter raised ImportError" assert self.SECRET not in json.dumps(row, default=str) + @pytest.mark.parametrize( + ("debug", "render_error"), + [ + pytest.param(False, RuntimeError, id="debug_off"), + pytest.param(True, RuntimeError, id="debug_on_exception"), + pytest.param(True, asyncio.CancelledError, id="debug_on_cancelled"), + pytest.param( + True, + _NeitherInterruptNorCancellation, + id="debug_on_other_base_exception", + ), + ], + ) + async def test_unrenderable_traceback_never_affects_the_row( + self, + debug, + render_error, + monkeypatch, + mock_write_client, + invocation_context, + dummy_arrow_schema, + ): + """A traceback that cannot be rendered falls back to a constant line. + + Uses a stock StreamHandler with logging.raiseExceptions on, Python's + default, whose handleError re-renders a failure's exception chain. + """ + monkeypatch.setattr(logging, "raiseExceptions", True) + tripwire = _Tripwire(render_error) + + def formatter(content, event_type): + failure = _unrenderable_exception(tripwire) + tripwire.armed = True + raise failure + + config = bigquery_agent_analytics_plugin.BigQueryLoggerConfig( + content_formatter=formatter, debug_content_formatter_errors=debug + ) + stream = io.StringIO() + + try: + with self._standard_handler_on_plugin_logger(stream): + row, drop_stats = await self._log_user_message( + config, mock_write_client, invocation_context, dummy_arrow_schema + ) + finally: + tripwire.armed = False + + assert row["error_message"] == ( + "content_formatter raised " + ) + assert ( + row["content"] + == bigquery_agent_analytics_plugin._FORMATTER_FAILED_SENTINEL + ) + assert drop_stats.get("formatter_failed") == 1 + logged = stream.getvalue() + assert logged.startswith(self.DEFAULT_WARNING) + assert ("[traceback could not be rendered]" in logged) is debug + assert "TRIPWIRE" not in logged + + @pytest.mark.parametrize( + "signal", [KeyboardInterrupt, SystemExit], ids=["interrupt", "exit"] + ) + async def test_debug_rendering_lets_interrupts_and_exits_propagate( + self, signal, mock_write_client, invocation_context, dummy_arrow_schema + ): + """KeyboardInterrupt and SystemExit raised while rendering propagate. + + A signal handler can deliver either at any bytecode, so one raised while + rendering may be genuine, and the plugin never swallows them. Anything + else raised there comes from the exception being rendered. + """ + + tripwire = _Tripwire(signal) + + def formatter(content, event_type): + failure = _unrenderable_exception(tripwire) + tripwire.armed = True + raise failure + + config = bigquery_agent_analytics_plugin.BigQueryLoggerConfig( + content_formatter=formatter, debug_content_formatter_errors=True + ) + + try: + with pytest.raises(signal, match="TRIPWIRE"): + await self._log_user_message( + config, mock_write_client, invocation_context, dummy_arrow_schema + ) + finally: + tripwire.armed = False + + @pytest.mark.parametrize( + "debug", [False, True], ids=["debug_off", "debug_on"] + ) + async def test_failing_log_handler_cannot_reach_the_formatter_exception( + self, + debug, + monkeypatch, + mock_write_client, + invocation_context, + dummy_arrow_schema, + ): + """A handler failure while warning never re-renders the formatter error. + + logging's handleError prints the failing handler's exception chain. The + warning is emitted after the formatter's exception is no longer being + handled, so that chain cannot reach it and run its hooks. + """ + monkeypatch.setattr(logging, "raiseExceptions", True) + tripwire = _Tripwire(RuntimeError) + + def formatter(content, event_type): + failure = _unrenderable_exception(tripwire) + tripwire.armed = True + raise failure + + config = bigquery_agent_analytics_plugin.BigQueryLoggerConfig( + content_formatter=formatter, debug_content_formatter_errors=debug + ) + closed_stream = io.StringIO() + closed_stream.close() + + try: + with self._standard_handler_on_plugin_logger(closed_stream): + row, drop_stats = await self._log_user_message( + config, mock_write_client, invocation_context, dummy_arrow_schema + ) + finally: + tripwire.armed = False + + assert ( + row["content"] + == bigquery_agent_analytics_plugin._FORMATTER_FAILED_SENTINEL + ) + assert drop_stats.get("formatter_failed") == 1 + class TestSafeCallbackDecorator: """Tests that _safe_callback prevents plugin errors from propagating.""" From 974e4642d6ac38d3b31067cb10ab83e9367e8e4e Mon Sep 17 00:00:00 2001 From: Haiyuan Cao Date: Mon, 28 Sep 2026 22:21:26 -0700 Subject: [PATCH 04/29] fix(plugins): keep formatter diagnosis from ever dropping the row Review of the previous commit showed three more ways that diagnosing a content_formatter failure could escape and drop the row, even after the hook-by-hook fixes: - a rendering hook raising SystemExit or KeyboardInterrupt, which the previous policy re-raised as possibly genuine; - type's own descriptors raising TypeError once a meta-metaclass drops `type` from the failed class's metaclass MRO, which disproved "classification cannot raise"; - a log handler or filter raising while the warning is emitted, which also misattributed an unsupported result as "raised RuntimeError". Instead of patching one hook at a time, all of diagnosis (naming the failed class, rendering the debug traceback, and emitting the warning) now runs in _diagnose_formatter_failure, behind one boundary that contains whatever is raised, of any type, and falls back to a constant note. It runs only after the fail-closed state is settled: the formatter's try/except now only records the failure and writes the sentinel, and the failure is counted before diagnosis starts. Nothing diagnosis does can reach the row, the sentinel, the counter, or the callback's caller. BaseException is contained deliberately. An interrupt raised by a hook, handler, or filter cannot be told apart from one a signal handler delivered, and letting it through would let the content under redaction abort the agent run and lose the row. Diagnosis never awaits, so a real asyncio cancellation is never swallowed; a signal that lands inside the short window is absorbed, and the next is delivered normally. Interrupts raised by the formatter call itself still propagate, as before. Two inner guards remain where they keep the warning: labeling falls back to "" when type's descriptors raise, and rendering falls back to a placeholder. Diagnosis still runs after the except block, so a failing handler's handleError cannot print the formatter's exception. A result whose __class__ merely claims to be str is now reported as an unsupported result instead of "raised TypeError". Tests: a property test injects RuntimeError, CancelledError, KeyboardInterrupt, SystemExit, and another BaseException at each diagnostic step (labeling, rendering, a handler, a filter) on both paths and requires the sentinel row and its count. Also added: the meta-metaclass regression on both paths, rendering hooks raising interrupts, a stderr check that a failing handler never prints the formatter's exception, and a guard that interrupts from the formatter call still propagate. Removing the boundary or any guard fails a test. Refs: https://github.com/GoogleCloudPlatform/BigQuery-Agent-Analytics-SDK/issues/485 (item A) Co-Authored-By: Claude Opus 5.5 (1M context) --- .../bigquery_agent_analytics_plugin.py | 197 ++++++---- .../test_bigquery_agent_analytics_plugin.py | 358 +++++++++++++++--- 2 files changed, 436 insertions(+), 119 deletions(-) diff --git a/src/google/adk/plugins/bigquery_agent_analytics_plugin.py b/src/google/adk/plugins/bigquery_agent_analytics_plugin.py index 97938f93105..478ef92140f 100644 --- a/src/google/adk/plugins/bigquery_agent_analytics_plugin.py +++ b/src/google/adk/plugins/bigquery_agent_analytics_plugin.py @@ -1068,8 +1068,7 @@ def _formatter_failure_message(cls: type, *, raised: bool) -> str: The class is named only by a trusted label, and a class without one by its nearest ancestor that has one. The exception's message, args, and traceback are never used: they can embed the content the formatter was protecting. - Nothing here runs code the class controls, so describing a failure cannot - raise, whatever the failure's class does. + Only type's own descriptors are read, so no code the class controls runs. Args: cls: The class of the exception the formatter raised, or of the value it @@ -1078,31 +1077,34 @@ def _formatter_failure_message(cls: type, *, raised: bool) -> str: Returns: A fixed-shape message such as ``content_formatter raised ImportError`` or - ``content_formatter returned unsupported type ``. + ``content_formatter returned unsupported type ``, + with ```` in place of the label when the class cannot be + read at all. """ outcome = "raised" if raised else "returned unsupported type" - for depth, ancestor in enumerate(_TYPE_MRO.__get__(cls)): - label = _trusted_class_label(ancestor) - if label is not None: - if depth == 0: - return f"content_formatter {outcome} {label}" - return f"content_formatter {outcome} " - # Unreachable for an instantiable class, whose MRO ends with object. + try: + for depth, ancestor in enumerate(_TYPE_MRO.__get__(cls)): + label = _trusted_class_label(ancestor) + if label is not None: + if depth == 0: + return f"content_formatter {outcome} {label}" + return f"content_formatter {outcome} " + except Exception: + # type's descriptors run no hooks but can still raise: they first check + # that the metaclass is a subtype of type by walking the metaclass's own + # MRO, which a meta-metaclass can rewrite after the class exists. + pass return f"content_formatter {outcome} " def _render_formatter_traceback(error: BaseException) -> str: """Renders a content_formatter exception's traceback for debug logging. - Rendering runs the exception's own code: its ``__str__`` and the attribute - hooks that expose its traceback and chained exceptions. Whatever that code - raises is contained here and replaced by a constant, so rendering can never - affect the row, with two exceptions. KeyboardInterrupt and SystemExit - propagate, as they do everywhere else in the plugin, because a signal - handler can deliver either at any bytecode, so one raised here may be - genuine. Any other BaseException, asyncio.CancelledError included, can only - come from the exception being rendered: rendering never awaits, so it cannot - receive a real cancellation. + Rendering runs code the exception's class controls: its ``__str__`` and the + attribute hooks that expose its traceback and chained exceptions. Whatever + that code raises, of any type, yields a constant placeholder instead, for + the reasons given in ``_diagnose_formatter_failure``, so the warning is + still logged. Args: error: The exception the formatter raised. @@ -1112,12 +1114,79 @@ def _render_formatter_traceback(error: BaseException) -> str: """ try: return "".join(traceback_module.format_exception(error)).rstrip("\n") - except (KeyboardInterrupt, SystemExit): - raise except BaseException: return "[traceback could not be rendered]" +def _diagnose_formatter_failure( + failed_type: type, + failure: Optional[BaseException], + *, + event_type: str, + debug: bool, +) -> str: + """Describes a content_formatter failure and logs its warning, best effort. + + Diagnosis runs only after the failure is fully handled: the row's content + is already the sentinel and the failure is already counted. Everything it + does, naming the failed class, rendering the debug traceback, and emitting + the warning through whatever filters and handlers are configured, sits + behind this one boundary. Whatever any of it raises, of any type, is + contained here and the constant fallback note is returned, so diagnosis + can never drop the row, change the sentinel or the counter, or escape the + callback. + + BaseException is contained on purpose. Rendering runs code the failed class + controls, and filters and handlers run arbitrary code. A KeyboardInterrupt + or SystemExit raised by any of them cannot be told apart from one that a + signal handler delivered, and letting it through would let the content + under redaction abort the agent run and lose the row. Diagnosis never + awaits, so it cannot swallow a real asyncio cancellation; a signal that + lands inside this short window is absorbed, and the next is delivered + normally. Interrupts raised by the formatter call itself still propagate. + + Args: + failed_type: The class of the exception the formatter raised, or of the + value it returned. + failure: The exception the formatter raised, or None when it returned an + unsupported value. + event_type: The type of the event being logged. + debug: Whether to append the rendered traceback to the warning. + + Returns: + The note for the error_message column. + """ + outcome = "raised" if failure is not None else "returned unsupported type" + note = f"content_formatter {outcome} " + try: + note = _formatter_failure_message(failed_type, raised=failure is not None) + if failure is None: + logger.warning( + "Content formatter returned an unsupported result type for" + " event %s; writing sentinel instead of original content.", + event_type, + ) + elif debug: + logger.warning( + "Content formatter failed for event %s; writing sentinel" + " instead of original content. Debug traceback:\n%s", + event_type, + _render_formatter_traceback(failure), + ) + else: + logger.warning( + "Content formatter failed for event %s; writing sentinel" + " instead of original content.", + event_type, + ) + except BaseException: + # Contained whatever it is, for the reasons in the docstring. It is not + # reported either: the logger may be what failed, and the exception can + # carry the content. + pass + return note + + # Recursion bound for _recursive_smart_truncate: id()-based cycle detection # cannot catch graphs that create new objects per access (Mock-like duck # typing); the cap turns unbounded recursion into a redacted leaf. @@ -2216,9 +2285,13 @@ class BigQueryLoggerConfig: pydantic ``BaseModel``, and ``google.api_core`` ``GoogleAPICallError``) are named. Any other class, including one your own code defines, is described by its nearest named ancestor, for example - ``content_formatter raised ``. An event that - already carries an ``error_message``, such as a ``TOOL_ERROR``, keeps - it first, followed by ``; `` and the formatter failure. + ``content_formatter raised ``, and a class + that cannot be read at all as ````. Describing the + failure and logging its warning are best effort: whatever they + raise is contained, so they never drop the row or change the + sentinel or the counter. An event that already carries an + ``error_message``, such as a ``TOOL_ERROR``, keeps it first, followed + by ``; `` and the formatter failure. gcs_bucket_name: GCS bucket for offloading large content. connection_id: BigQuery connection ID for ObjectRef columns. log_session_metadata: Whether to log session metadata. @@ -2280,9 +2353,9 @@ class BigQueryLoggerConfig: has configured: the console, the log file that ``adk run`` writes, and anything that forwards logs elsewhere, such as a managed runtime shipping stderr to Cloud Logging. Enable it only where that content - may be seen. Rendering is best effort: if the exception's own code - fails while it is rendered, a constant placeholder is logged instead, - and the row is unaffected. The plugin never writes the traceback to + may be seen. Rendering is best effort: whatever the exception's own + code raises while it is rendered, a constant placeholder is logged + instead, and the row is unaffected. The plugin never writes the traceback to BigQuery; the row's ``error_message`` still names only the exception class. ``False`` (the default) logs a constant message with no traceback. @@ -7327,12 +7400,15 @@ async def _log_event( timestamp = datetime.now(timezone.utc) formatter_error: Optional[str] = None - formatter_raised = False - formatter_traceback: Optional[str] = None if self.config.content_formatter: + failed_type: Optional[type] = None + failure: Optional[Exception] = None try: formatted = self.config.content_formatter(raw_content, event_type) - if isinstance(formatted, str): + # The real type, not isinstance: an object whose __class__ claims to + # be str is not one, and normalizing it would raise as if the + # formatter had. + if issubclass(type(formatted), str): if type(formatted) is not str: # Normalize str subclasses to the exact built-in. formatted = str.__str__(formatted) @@ -7359,53 +7435,38 @@ async def _log_event( # the parser's unconditional str(content) fallback OUTSIDE this # fail-closed boundary, where a payload-controlled __str__ can # republish the original content or raise into the safe - # callback's traceback log. The - # message is CONSTANT: even a class NAME can be payload-derived - # via type(name, ...), so error_message uses only a trusted class - # label. - logger.warning( - "Content formatter returned an unsupported result type for" - " event %s; writing sentinel instead of original content.", - event_type, - ) - formatter_error = _formatter_failure_message( - type(formatted), raised=False - ) + # callback's traceback log. + failed_type = type(formatted) formatted = _FORMATTER_FAILED_SENTINEL - self._count_local_drop("formatter_failed") raw_content = formatted except Exception as e: # Fail CLOSED: the formatter is a redaction/privacy # boundary, so its failure must never fall back to the unformatted # payload. The exception message and traceback can embed the # protected content, and even the class NAME can be payload-derived - # via type(name, ...), so error_message uses only a trusted class - # label and the warning below stays CONSTANT. The traceback is - # rendered, best effort, only when debug_content_formatter_errors - # opts in; no handler ever receives the live exception. - formatter_raised = True - formatter_error = _formatter_failure_message(type(e), raised=True) - if self.config.debug_content_formatter_errors: - formatter_traceback = _render_formatter_traceback(e) + # via type(name, ...), so diagnosis below names only a trusted class + # label, keeps the default warning CONSTANT, and renders the + # traceback only when debug_content_formatter_errors opts in. + failed_type = type(e) + failure = e raw_content = _FORMATTER_FAILED_SENTINEL + if failed_type is not None: + # The sentinel is in place and the failure is counted before any + # diagnosis runs, and diagnosis contains whatever it raises, so + # describing the failure can never drop the row or change either. + # It runs after the except block so that a failing log handler, + # whose handleError prints the exception being handled, cannot + # reach the formatter's exception. self._count_local_drop("formatter_failed") - if formatter_raised: - # Logged only once the formatter's exception is no longer being - # handled: a failing handler's handleError prints its own exception - # chain, which would otherwise reach that exception and run its code. - if formatter_traceback is None: - logger.warning( - "Content formatter failed for event %s; writing sentinel" - " instead of original content.", - event_type, - ) - else: - logger.warning( - "Content formatter failed for event %s; writing sentinel" - " instead of original content. Debug traceback:\n%s", - event_type, - formatter_traceback, - ) + formatter_error = _diagnose_formatter_failure( + failed_type, + failure, + event_type=event_type, + debug=self.config.debug_content_formatter_errors, + ) + # The except clause would have dropped this reference itself: the + # exception's traceback holds this frame, which holds the exception. + failure = None # The event's own diagnostic (e.g. a TOOL_ERROR's exception text) stays # first and intact so an error row keeps its primary cause; a formatter diff --git a/tests/unittests/plugins/test_bigquery_agent_analytics_plugin.py b/tests/unittests/plugins/test_bigquery_agent_analytics_plugin.py index 9ba6e882286..c2a70f3b3bd 100644 --- a/tests/unittests/plugins/test_bigquery_agent_analytics_plugin.py +++ b/tests/unittests/plugins/test_bigquery_agent_analytics_plugin.py @@ -3919,6 +3919,37 @@ def __str__(self): return _UnrenderableError() +class _MetaclassMroBreaker: + """Builds a metaclass whose own MRO can drop `type` after classes exist. + + Reading a class through type's descriptors first checks that the class's + metaclass is a subtype of type, by walking the metaclass's MRO, so every + such read raises TypeError once the MRO is broken. pytest reads the same + descriptors when it reports a failure, so tests restore the MRO before + anything is reported. + """ + + def __init__(self): + self.broken = False + breaker = self + + class _MetaMeta(type): + + def mro(cls): + return [cls, object] if breaker.broken else super().mro() + + self.metaclass = _MetaMeta("_BreakableMeta", (type,), {}) + + def break_mro(self): + self.broken = True + self.metaclass.__bases__ = (type,) # Recomputes the metaclass MRO. + + def restore(self): + if self.broken: + self.broken = False + self.metaclass.__bases__ = (type,) + + def _raise_import_error(content, event_type): raise ImportError(f"cannot import name 'redact' (formatting {content})") @@ -3990,6 +4021,16 @@ class LocalPart(types.Part): return LocalPart() +def _return_object_claiming_to_be_str(content, event_type): + class ClaimsToBeStr: + + @property + def __class__(self): + return str + + return ClaimsToBeStr() + + def _return_pydantic_model(content, event_type): class RedactedPayload(BaseModel): text: str = "[REDACTED]" @@ -4019,6 +4060,10 @@ class TestContentFormatterFailureDiagnostics: "Content formatter failed for event USER_MESSAGE_RECEIVED; writing" " sentinel instead of original content." ) + UNSUPPORTED_WARNING = ( + "Content formatter returned an unsupported result type for event" + " USER_MESSAGE_RECEIVED; writing sentinel instead of original content." + ) @pytest.fixture(autouse=True) def _unbind_registered_payload_classes(self): @@ -4045,12 +4090,45 @@ async def _log_user_message( ) return row, plugin.get_drop_stats() + async def _log_user_message_contained( + self, + config, + mock_write_client, + invocation_context, + dummy_arrow_schema, + *, + cleanup=None, + ): + """Like _log_user_message, but anything escaping becomes a test failure. + + An escaping KeyboardInterrupt would otherwise stop the whole test run. + cleanup runs before pytest reports anything, so that hostile hooks can + be disarmed first. + """ + escaped = None + try: + return await self._log_user_message( + config, mock_write_client, invocation_context, dummy_arrow_schema + ) + except BaseException as error: # pylint: disable=broad-exception-caught + escaped = error + finally: + if cleanup is not None: + cleanup() + if type(escaped) is AssertionError: + # Raised by _get_captured_event_dict_async: the callback returned, but + # no row reached the write path. + pytest.fail(f"no row was written: {escaped}", pytrace=False) + pytest.fail( + f"{type(escaped).__name__} escaped the plugin callback", pytrace=False + ) + @staticmethod def _formatter_warnings(caplog): return [ record for record in caplog.records - if record.getMessage().startswith("Content formatter failed") + if record.getMessage().startswith("Content formatter ") ] @staticmethod @@ -4146,6 +4224,12 @@ def _standard_handler_on_plugin_logger(stream): "content_formatter returned unsupported type ", id="part_subclass_result", ), + pytest.param( + _return_object_claiming_to_be_str, + "content_formatter returned unsupported type" + " ", + id="object_claiming_to_be_str_result", + ), pytest.param( _return_pydantic_model, "content_formatter returned unsupported type" @@ -4403,6 +4487,8 @@ async def test_debug_flag_logs_traceback_locally_but_not_to_the_row( pytest.param(False, RuntimeError, id="debug_off"), pytest.param(True, RuntimeError, id="debug_on_exception"), pytest.param(True, asyncio.CancelledError, id="debug_on_cancelled"), + pytest.param(True, KeyboardInterrupt, id="debug_on_interrupt"), + pytest.param(True, SystemExit, id="debug_on_exit"), pytest.param( True, _NeitherInterruptNorCancellation, @@ -4421,8 +4507,10 @@ async def test_unrenderable_traceback_never_affects_the_row( ): """A traceback that cannot be rendered falls back to a constant line. - Uses a stock StreamHandler with logging.raiseExceptions on, Python's - default, whose handleError re-renders a failure's exception chain. + Whatever the exception's own code raises while it is rendered, including + KeyboardInterrupt and SystemExit, is contained: the row is written and + counted, and the warning is still logged with a placeholder. Uses a stock + StreamHandler with logging.raiseExceptions on, Python's default. """ monkeypatch.setattr(logging, "raiseExceptions", True) tripwire = _Tripwire(render_error) @@ -4437,14 +4525,18 @@ def formatter(content, event_type): ) stream = io.StringIO() - try: - with self._standard_handler_on_plugin_logger(stream): - row, drop_stats = await self._log_user_message( - config, mock_write_client, invocation_context, dummy_arrow_schema - ) - finally: + def disarm(): tripwire.armed = False + with self._standard_handler_on_plugin_logger(stream): + row, drop_stats = await self._log_user_message_contained( + config, + mock_write_client, + invocation_context, + dummy_arrow_schema, + cleanup=disarm, + ) + assert row["error_message"] == ( "content_formatter raised " ) @@ -4459,81 +4551,245 @@ def formatter(content, event_type): assert "TRIPWIRE" not in logged @pytest.mark.parametrize( - "signal", [KeyboardInterrupt, SystemExit], ids=["interrupt", "exit"] + "interrupt", + [KeyboardInterrupt, SystemExit, asyncio.CancelledError], + ids=["keyboard_interrupt", "system_exit", "cancelled_error"], ) - async def test_debug_rendering_lets_interrupts_and_exits_propagate( - self, signal, mock_write_client, invocation_context, dummy_arrow_schema + async def test_interrupts_raised_by_the_formatter_call_still_propagate( + self, interrupt, mock_write_client, invocation_context, dummy_arrow_schema ): - """KeyboardInterrupt and SystemExit raised while rendering propagate. + """Containment covers diagnosis only, never the formatter call itself. - A signal handler can deliver either at any bytecode, so one raised while - rendering may be genuine, and the plugin never swallows them. Anything - else raised there comes from the exception being rendered. + The fail-closed boundary around the call catches Exception, so an + interrupt or cancellation raised while the formatter runs reaches the + caller exactly as before. """ - tripwire = _Tripwire(signal) - def formatter(content, event_type): - failure = _unrenderable_exception(tripwire) - tripwire.armed = True - raise failure + raise interrupt("raised by the formatter call") config = bigquery_agent_analytics_plugin.BigQueryLoggerConfig( content_formatter=formatter, debug_content_formatter_errors=True ) - try: - with pytest.raises(signal, match="TRIPWIRE"): - await self._log_user_message( - config, mock_write_client, invocation_context, dummy_arrow_schema - ) - finally: - tripwire.armed = False + with pytest.raises(interrupt, match="raised by the formatter call"): + await self._log_user_message( + config, mock_write_client, invocation_context, dummy_arrow_schema + ) - @pytest.mark.parametrize( - "debug", [False, True], ids=["debug_off", "debug_on"] - ) - async def test_failing_log_handler_cannot_reach_the_formatter_exception( + async def test_failing_log_handler_never_prints_the_formatter_exception( self, - debug, + capsys, monkeypatch, mock_write_client, invocation_context, dummy_arrow_schema, ): - """A handler failure while warning never re-renders the formatter error. + """A handler that fails while warning cannot print the formatter error. - logging's handleError prints the failing handler's exception chain. The - warning is emitted after the formatter's exception is no longer being - handled, so that chain cannot reach it and run its hooks. + logging's handleError prints the failing handler's exception chain to + stderr. The warning is emitted after the formatter's exception is no + longer being handled, so that chain never includes it or the content its + message embeds. """ monkeypatch.setattr(logging, "raiseExceptions", True) - tripwire = _Tripwire(RuntimeError) - - def formatter(content, event_type): - failure = _unrenderable_exception(tripwire) - tripwire.armed = True - raise failure - config = bigquery_agent_analytics_plugin.BigQueryLoggerConfig( - content_formatter=formatter, debug_content_formatter_errors=debug + content_formatter=_raise_import_error ) closed_stream = io.StringIO() closed_stream.close() - try: - with self._standard_handler_on_plugin_logger(closed_stream): - row, drop_stats = await self._log_user_message( - config, mock_write_client, invocation_context, dummy_arrow_schema + with self._standard_handler_on_plugin_logger(closed_stream): + row, drop_stats = await self._log_user_message_contained( + config, mock_write_client, invocation_context, dummy_arrow_schema + ) + + assert row["error_message"] == "content_formatter raised ImportError" + assert drop_stats.get("formatter_failed") == 1 + stderr = capsys.readouterr().err + assert "--- Logging error ---" in stderr + assert self.SECRET not in stderr + + @pytest.mark.parametrize( + ("raised", "expected_error_message", "expected_warning"), + [ + pytest.param( + True, + "content_formatter raised ", + DEFAULT_WARNING, + id="raised", + ), + pytest.param( + False, + "content_formatter returned unsupported type ", + UNSUPPORTED_WARNING, + id="returned", + ), + ], + ) + async def test_class_that_cannot_be_read_still_gets_a_row_and_a_warning( + self, + raised, + expected_error_message, + expected_warning, + mock_write_client, + invocation_context, + dummy_arrow_schema, + caplog, + ): + """A failed class that even type's own descriptors reject is not named. + + A meta-metaclass can drop type from the failed class's metaclass MRO + after the class exists, so every descriptor read raises TypeError. The + row is still written and counted, the constant warning is still logged, + and nothing from the formatter's exception reaches either. + """ + breaker = _MetaclassMroBreaker() + + def formatter(content, event_type): + name = _identifier_from_content(content) + if raised: + failure = breaker.metaclass(name, (ValueError,), {})( + f"cannot redact {content}" ) - finally: - tripwire.armed = False + else: + failure = breaker.metaclass(name, (), {})() + breaker.break_mro() + if raised: + raise failure + return failure + + config = bigquery_agent_analytics_plugin.BigQueryLoggerConfig( + content_formatter=formatter + ) + + with caplog.at_level(logging.WARNING): + row, drop_stats = await self._log_user_message_contained( + config, + mock_write_client, + invocation_context, + dummy_arrow_schema, + cleanup=breaker.restore, + ) + assert row["error_message"] == expected_error_message assert ( row["content"] == bigquery_agent_analytics_plugin._FORMATTER_FAILED_SENTINEL ) assert drop_stats.get("formatter_failed") == 1 + assert [ + record.getMessage() for record in self._formatter_warnings(caplog) + ] == [expected_warning] + for payload_text in (self.SECRET, self.PAYLOAD_IDENTIFIER): + assert payload_text not in json.dumps(row, default=str) + assert payload_text not in caplog.text + + @pytest.mark.parametrize( + "injected", + [ + RuntimeError, + asyncio.CancelledError, + KeyboardInterrupt, + SystemExit, + _NeitherInterruptNorCancellation, + ], + ids=[ + "runtime_error", + "cancelled_error", + "keyboard_interrupt", + "system_exit", + "other_base_exception", + ], + ) + @pytest.mark.parametrize( + ("step", "raised"), + [ + pytest.param("label", True, id="label-raised"), + pytest.param("label", False, id="label-returned"), + pytest.param("render", True, id="render-raised"), + pytest.param("log_handler", True, id="log_handler-raised"), + pytest.param("log_handler", False, id="log_handler-returned"), + pytest.param("log_filter", True, id="log_filter-raised"), + pytest.param("log_filter", False, id="log_filter-returned"), + ], + ) + async def test_a_raise_anywhere_in_diagnosis_leaves_the_row_intact( + self, + step, + raised, + injected, + mock_write_client, + invocation_context, + dummy_arrow_schema, + caplog, + ): + """Whatever any diagnostic step raises, the row is written unchanged. + + Diagnosis (naming the failed class, rendering the debug traceback, and + emitting the warning through the logger's filters and handlers) runs + after the failure is handled, behind one boundary. Each step here raises + each kind of exception, and the sentinel row, its drop count, and a + payload-free error_message must still come out, with nothing escaping. + """ + plugin_module = bigquery_agent_analytics_plugin + plugin_logger = logging.getLogger("google_adk." + plugin_module.__name__) + + def raise_injected(*args, **kwargs): + raise injected(f"injected into {step}") + + def raise_for_formatter_warnings(record): + if record.getMessage().startswith("Content formatter "): + raise_injected() + return True + + class _RaisingHandler(logging.Handler): + + def emit(self, record): + raise_for_formatter_warnings(record) + + config = plugin_module.BigQueryLoggerConfig( + content_formatter=_raise_import_error if raised else _return_tuple, + debug_content_formatter_errors=True, + ) + handler = _RaisingHandler() + injections = { + "label": mock.patch.object( + plugin_module, "_formatter_failure_message", raise_injected + ), + "render": mock.patch.object( + plugin_module, "_render_formatter_traceback", raise_injected + ), + "log_handler": contextlib.nullcontext(), + "log_filter": contextlib.nullcontext(), + } + if step == "log_handler": + plugin_logger.addHandler(handler) + if step == "log_filter": + plugin_logger.addFilter(raise_for_formatter_warnings) + try: + with injections[step], caplog.at_level(logging.WARNING): + row, drop_stats = await self._log_user_message_contained( + config, mock_write_client, invocation_context, dummy_arrow_schema + ) + finally: + plugin_logger.removeHandler(handler) + plugin_logger.removeFilter(raise_for_formatter_warnings) + + outcome = "raised" if raised else "returned unsupported type" + if step == "label": + expected_error_message = f"content_formatter {outcome} " + elif raised: + expected_error_message = "content_formatter raised ImportError" + else: + expected_error_message = ( + "content_formatter returned unsupported type tuple" + ) + assert row["error_message"] == expected_error_message + assert row["content"] == plugin_module._FORMATTER_FAILED_SENTINEL + assert drop_stats.get("formatter_failed") == 1 + assert self.SECRET not in json.dumps(row, default=str) + assert self.SECRET not in caplog.text class TestSafeCallbackDecorator: From 7fae96552a4f032e7d5ca0e25c3ee9a0e315ea7e Mon Sep 17 00:00:00 2001 From: Haiyuan Cao Date: Mon, 28 Sep 2026 23:39:57 -0700 Subject: [PATCH 05/29] fix(plugins): close remaining formatter paths that leak or drop rows Review of the previous commit found three more paths, each also present on main, by which a content_formatter failure could still leak payload text or lose the row: - The failure warning can be logged while the caller is handling an exception, as ADK is when it runs error callbacks. A log handler that fails prints the exception being handled through handleError, so a closed stream could print the caller's exception, or the formatter's when it re-raised that one, to stderr. - The result check used isinstance(formatted, (dict, list)), which falls back to the object's own __class__. A returned object could raise CancelledError, SystemExit, or KeyboardInterrupt there and lose the row, or raise an ordinary exception that was then reported as the formatter having raised. - A rejected coroutine was released unstarted, so Python warned "coroutine '' was never awaited", and the formatter can set that name from the content. Everything after the formatter call now runs in _settle_formatter_outcome, behind the one boundary that already contained diagnosis: judging the result, closing a rejected coroutine or generator, naming the class, rendering the debug traceback, and logging. The result is judged only by its real type, through issubclass(type(x), ...) and identity, which runs none of its code. A str, dict, or list subclass is still accepted; an object whose __class__ merely claims to be one is now a counted formatter failure rather than an [UNSUPPORTED_OBJECT] row. A rejected coroutine or generator is closed before release, which for an unstarted one runs no code and emits no warning. The warning is emitted while a constant stand-in exception, raised from None, is being handled, so a failing handler prints only that. The formatter call itself still lets KeyboardInterrupt, SystemExit, and CancelledError propagate. Docs: the content_formatter docstring now says "raises an Exception", explains judging by real type, and notes that interrupts from the call propagate. The boundary's docstring states plainly that a genuine signal delivered while it runs, including while a handler blocks, is absorbed. The debug flag notes that a failing handler echoes the rendered traceback to stderr. Tests: a caller handling a payload exception with a closed handler (re-raised and distinct); results whose __class__ property or __getattribute__ raises each interrupt type or claims to be a dict; an async formatter and a payload-named coroutine, with warnings recorded; and two more steps in the injection property test, judging the result and closing a rejected generator. Removing any guard fails a test. Refs: https://github.com/GoogleCloudPlatform/BigQuery-Agent-Analytics-SDK/issues/485 (item A) Co-Authored-By: Claude Opus 5.5 (1M context) --- .../bigquery_agent_analytics_plugin.py | 281 ++++++++++-------- .../test_bigquery_agent_analytics_plugin.py | 246 ++++++++++++++- 2 files changed, 397 insertions(+), 130 deletions(-) diff --git a/src/google/adk/plugins/bigquery_agent_analytics_plugin.py b/src/google/adk/plugins/bigquery_agent_analytics_plugin.py index 478ef92140f..373af9abe98 100644 --- a/src/google/adk/plugins/bigquery_agent_analytics_plugin.py +++ b/src/google/adk/plugins/bigquery_agent_analytics_plugin.py @@ -51,6 +51,8 @@ import re import threading import time +from types import CoroutineType +from types import GeneratorType from types import MappingProxyType from types import TracebackType from typing import Any @@ -1103,7 +1105,7 @@ def _render_formatter_traceback(error: BaseException) -> str: Rendering runs code the exception's class controls: its ``__str__`` and the attribute hooks that expose its traceback and chained exceptions. Whatever that code raises, of any type, yields a constant placeholder instead, for - the reasons given in ``_diagnose_formatter_failure``, so the warning is + the reasons given in ``_settle_formatter_outcome``, so the warning is still logged. Args: @@ -1118,73 +1120,138 @@ def _render_formatter_traceback(error: BaseException) -> str: return "[traceback could not be rendered]" -def _diagnose_formatter_failure( - failed_type: type, - failure: Optional[BaseException], +class _WarningIsolation(Exception): + """The exception being handled while the formatter-failure warning is logged. + + A log handler that fails calls ``handleError``, which prints the exception + being handled. Without this stand-in that could be the formatter's + exception, or one the caller is handling, and either can carry the content. + """ + + +def _natively_parsed(formatted: Any, result_type: type) -> bool: + """Whether the parser logs a formatter result of this real type natively. + + Identity and conditional formatters legitimately return these shapes. Model + shapes must be the EXACT class, compared by identity: a subclass can + override an attribute the parser reads, and an equality check would run the + result class's metaclass. str, dict, and list subclasses are admitted, + because the parser routes them through its hardened recursive sanitizer. + Only ``result_type``, the result's real type, is consulted: ``isinstance`` + would fall back to the object's own ``__class__`` and run its code. + + Args: + formatted: What the formatter returned. + result_type: ``type(formatted)``. + + Returns: + Whether ``formatted`` can be logged as it is. + """ + return ( + formatted is None + or issubclass(result_type, (str, dict, list)) + or result_type is types.Content + or result_type is types.Part + or result_type is LlmRequest + ) + + +def _settle_formatter_outcome( + formatted: Any, + failure: Optional[Exception], *, event_type: str, debug: bool, -) -> str: - """Describes a content_formatter failure and logs its warning, best effort. - - Diagnosis runs only after the failure is fully handled: the row's content - is already the sentinel and the failure is already counted. Everything it - does, naming the failed class, rendering the debug traceback, and emitting - the warning through whatever filters and handlers are configured, sits - behind this one boundary. Whatever any of it raises, of any type, is - contained here and the constant fallback note is returned, so diagnosis - can never drop the row, change the sentinel or the counter, or escape the - callback. - - BaseException is contained on purpose. Rendering runs code the failed class - controls, and filters and handlers run arbitrary code. A KeyboardInterrupt - or SystemExit raised by any of them cannot be told apart from one that a - signal handler delivered, and letting it through would let the content - under redaction abort the agent run and lose the row. Diagnosis never - awaits, so it cannot swallow a real asyncio cancellation; a signal that - lands inside this short window is absorbed, and the next is delivered - normally. Interrupts raised by the formatter call itself still propagate. +) -> tuple[Any, Optional[str]]: + """Decides what the row logs after the content_formatter call, fail closed. + + Everything after the call happens behind this one boundary: + - judging a returned result by its real type; + - closing a rejected coroutine or generator; + - naming the failed class; + - rendering the debug traceback; + - emitting the warning through whatever filters and handlers are + configured. + + Whatever any step raises, of any type, is contained here. The outcome then + falls back to the sentinel and a constant note, so nothing after the call + can drop the row, leave the sentinel out, or escape the callback. + + BaseException is contained on purpose. These steps run code that the + failed class or the rejected result controls, and code in filters and + handlers. A KeyboardInterrupt or SystemExit raised there cannot be told + apart from one that a signal handler delivered, and letting it through + would let the content under redaction abort the agent run and lose the + row. The cost is that a signal delivered while these steps run is absorbed + and the process does not act on it. That includes a signal that arrives + while a configured handler blocks, for example on network I/O. An + orchestrator that follows SIGTERM with SIGKILL then loses the rows still + buffered. Nothing here awaits, so a real asyncio cancellation is never + absorbed, and interrupts raised by the formatter call itself still + propagate. Args: - failed_type: The class of the exception the formatter raised, or of the - value it returned. - failure: The exception the formatter raised, or None when it returned an - unsupported value. + formatted: What the formatter returned; ignored when ``failure`` is set. + failure: The exception the formatter raised, or None if it returned. event_type: The type of the event being logged. debug: Whether to append the rendered traceback to the warning. Returns: - The note for the error_message column. + ``(formatted, None)`` when the parser can log the result as it is. A str + subclass is normalized to the exact built-in. Otherwise + ``(_FORMATTER_FAILED_SENTINEL, note)``, where ``note`` is the text for + the error_message column. """ outcome = "raised" if failure is not None else "returned unsupported type" note = f"content_formatter {outcome} " try: - note = _formatter_failure_message(failed_type, raised=failure is not None) - if failure is None: - logger.warning( - "Content formatter returned an unsupported result type for" - " event %s; writing sentinel instead of original content.", - event_type, - ) - elif debug: - logger.warning( - "Content formatter failed for event %s; writing sentinel" - " instead of original content. Debug traceback:\n%s", - event_type, - _render_formatter_traceback(failure), - ) + if failure is not None: + failed_type: type = type(failure) else: - logger.warning( - "Content formatter failed for event %s; writing sentinel" - " instead of original content.", - event_type, - ) + failed_type = type(formatted) + if _natively_parsed(formatted, failed_type): + if failed_type is not str and issubclass(failed_type, str): + formatted = str.__str__(formatted) + return formatted, None + # A non-native result would reach the parser's str() fallback, where + # a payload-controlled __str__ can republish the content, so it is + # rejected. Close a coroutine first: released unstarted, it warns + # "coroutine '' was never awaited", and the formatter can set + # that name from the content. A generator is closed too, so that its + # cleanup code runs here, contained, rather than at collection. + if failed_type is CoroutineType: + CoroutineType.close(formatted) + elif failed_type is GeneratorType: + GeneratorType.close(formatted) + note = _formatter_failure_message(failed_type, raised=failure is not None) + try: + raise _WarningIsolation from None + except _WarningIsolation: + if failure is None: + logger.warning( + "Content formatter returned an unsupported result type for" + " event %s; writing sentinel instead of original content.", + event_type, + ) + elif debug: + logger.warning( + "Content formatter failed for event %s; writing sentinel" + " instead of original content. Debug traceback:\n%s", + event_type, + _render_formatter_traceback(failure), + ) + else: + logger.warning( + "Content formatter failed for event %s; writing sentinel" + " instead of original content.", + event_type, + ) except BaseException: # Contained whatever it is, for the reasons in the docstring. It is not # reported either: the logger may be what failed, and the exception can # carry the content. pass - return note + return _FORMATTER_FAILED_SENTINEL, note # Recursion bound for _recursive_smart_truncate: id()-based cycle detection @@ -2273,25 +2340,32 @@ class BigQueryLoggerConfig: content_formatter: Optional custom formatter for content, called as ``content_formatter(content, event_type)``. It is treated as a redaction boundary, so a failure never falls back to the original - content: if it raises, or returns anything other than a ``str``, - ``dict``, ``list``, ``None``, or an exact ``types.Content``, - ``types.Part``, or ``LlmRequest``, the row is written with content + content: if it raises an ``Exception``, or returns anything other + than ``None``, a ``str``, ``dict``, or ``list``, or an exact + ``types.Content``, ``types.Part``, or ``LlmRequest``, the row is + written with content ``[FORMATTER_FAILED]``, the ``formatter_failed`` counter of ``get_drop_stats()`` is incremented, and ``error_message`` names the failure by class only, for example ``content_formatter raised - ImportError``. Because a class can be created or renamed at runtime + ImportError``. A result is judged by its real type, so a subclass of + ``str``, ``dict``, or ``list`` is accepted and an object whose + ``__class__`` merely claims to be one is not. Because a class can be + created or renamed at runtime with a name taken from the content, only built-in types and a few trusted classes (``LlmRequest``, ``types.Content``, ``types.Part``, pydantic ``BaseModel``, and ``google.api_core`` ``GoogleAPICallError``) are named. Any other class, including one your own code defines, is described by its nearest named ancestor, for example ``content_formatter raised ``, and a class - that cannot be read at all as ````. Describing the - failure and logging its warning are best effort: whatever they - raise is contained, so they never drop the row or change the - sentinel or the counter. An event that already carries an - ``error_message``, such as a ``TOOL_ERROR``, keeps it first, followed - by ``; `` and the formatter failure. + that cannot be read at all as ````. Judging the + result, describing the failure, and logging its warning are best + effort: whatever they raise is contained, so they never drop the + row or change the sentinel or the counter. A ``KeyboardInterrupt``, + ``SystemExit``, or ``asyncio.CancelledError`` raised by the + formatter call itself still propagates, and no row is written. An + event that already carries an ``error_message``, such as a + ``TOOL_ERROR``, keeps it first, followed by ``; `` and the formatter + failure. gcs_bucket_name: GCS bucket for offloading large content. connection_id: BigQuery connection ID for ObjectRef columns. log_session_metadata: Whether to log session metadata. @@ -2353,12 +2427,14 @@ class BigQueryLoggerConfig: has configured: the console, the log file that ``adk run`` writes, and anything that forwards logs elsewhere, such as a managed runtime shipping stderr to Cloud Logging. Enable it only where that content - may be seen. Rendering is best effort: whatever the exception's own - code raises while it is rendered, a constant placeholder is logged - instead, and the row is unaffected. The plugin never writes the traceback to - BigQuery; the row's ``error_message`` still names only the exception - class. ``False`` (the default) logs a constant message with no - traceback. + may be seen. That includes stderr when a handler fails: logging's + ``handleError`` prints the failing record's arguments, and the + rendered traceback is one of them. Rendering is best effort: + whatever the exception's own code raises while it is rendered, a + constant placeholder is logged instead, and the row is unaffected. + The plugin never writes the traceback to BigQuery; the row's + ``error_message`` still names only the exception class. ``False`` + (the default) logs a constant message with no traceback. """ enabled: bool = True @@ -7401,72 +7477,35 @@ async def _log_event( timestamp = datetime.now(timezone.utc) formatter_error: Optional[str] = None if self.config.content_formatter: - failed_type: Optional[type] = None + formatted: Any = None failure: Optional[Exception] = None try: formatted = self.config.content_formatter(raw_content, event_type) - # The real type, not isinstance: an object whose __class__ claims to - # be str is not one, and normalizing it would raise as if the - # formatter had. - if issubclass(type(formatted), str): - if type(formatted) is not str: - # Normalize str subclasses to the exact built-in. - formatted = str.__str__(formatted) - elif formatted is not None and not ( - # Every shape the parser handles NATIVELY: identity and - # conditional formatters legitimately return these, and the - # Str/Content/None-only gate destroyed untransformed - # LlmRequest/dict/list events. - # Model shapes require the EXACT class: a subclass can - # override an attribute the parser reads OUTSIDE this - # boundary and raise a payload-bearing exception into the - # safe callback's traceback log. dict/list subclasses stay isinstance-based — the - # parser routes them through the hardened recursive - # sanitizer, whose protocol boundary already fails closed. - # Compared by identity: `in` would call the result class's - # metaclass __eq__, which can raise anything or claim a match. - any( - type(formatted) is shape - for shape in (types.Content, types.Part, LlmRequest) - ) - or isinstance(formatted, (dict, list)) - ): - # The formatter is typed Any: a non-native result would reach - # the parser's unconditional str(content) fallback OUTSIDE this - # fail-closed boundary, where a payload-controlled __str__ can - # republish the original content or raise into the safe - # callback's traceback log. - failed_type = type(formatted) - formatted = _FORMATTER_FAILED_SENTINEL - raw_content = formatted except Exception as e: # Fail CLOSED: the formatter is a redaction/privacy # boundary, so its failure must never fall back to the unformatted # payload. The exception message and traceback can embed the # protected content, and even the class NAME can be payload-derived - # via type(name, ...), so diagnosis below names only a trusted class - # label, keeps the default warning CONSTANT, and renders the - # traceback only when debug_content_formatter_errors opts in. - failed_type = type(e) + # via type(name, ...), so the failure is only ever named by a + # trusted class label, and its traceback is logged only when + # debug_content_formatter_errors opts in. failure = e - raw_content = _FORMATTER_FAILED_SENTINEL - if failed_type is not None: - # The sentinel is in place and the failure is counted before any - # diagnosis runs, and diagnosis contains whatever it raises, so - # describing the failure can never drop the row or change either. - # It runs after the except block so that a failing log handler, - # whose handleError prints the exception being handled, cannot - # reach the formatter's exception. + # Everything after the call runs behind _settle_formatter_outcome's + # one boundary, which contains whatever it raises: judging the result, + # closing a rejected coroutine, naming the class, and logging. It runs + # after the except block, so the formatter's exception is no longer + # the one being handled. + raw_content, formatter_error = _settle_formatter_outcome( + formatted, + failure, + event_type=event_type, + debug=self.config.debug_content_formatter_errors, + ) + if formatter_error is not None: self._count_local_drop("formatter_failed") - formatter_error = _diagnose_formatter_failure( - failed_type, - failure, - event_type=event_type, - debug=self.config.debug_content_formatter_errors, - ) - # The except clause would have dropped this reference itself: the - # exception's traceback holds this frame, which holds the exception. - failure = None + # The except clause would have dropped this reference itself: the + # exception's traceback holds this frame, which holds the exception. + formatted = failure = None # The event's own diagnostic (e.g. a TOOL_ERROR's exception text) stays # first and intact so an error row keeps its primary cause; a formatter diff --git a/tests/unittests/plugins/test_bigquery_agent_analytics_plugin.py b/tests/unittests/plugins/test_bigquery_agent_analytics_plugin.py index c2a70f3b3bd..36dff868632 100644 --- a/tests/unittests/plugins/test_bigquery_agent_analytics_plugin.py +++ b/tests/unittests/plugins/test_bigquery_agent_analytics_plugin.py @@ -17,6 +17,7 @@ import concurrent.futures import contextlib import dataclasses +import gc import io import json import logging @@ -26,6 +27,7 @@ import threading import time from unittest import mock +import warnings from google.adk.agents import base_agent from google.adk.agents.callback_context import CallbackContext @@ -3950,6 +3952,36 @@ def restore(self): self.metaclass.__bases__ = (type,) +def _result_with_class_hook(tripwire, via, claims=None): + """Returns an object whose own __class__ lookup fires or lies. + + isinstance falls back to an object's __class__ when its real type does not + match, which runs this code. `claims` is what the lookup reports instead of + the real class. + """ + if via == "property": + + class _Result: + + @property + def __class__(self): + tripwire.fire("__class__ property") + return claims if claims is not None else type(self) + + else: + + class _Result: + + def __getattribute__(self, name): + if name == "__class__": + tripwire.fire("__getattribute__('__class__')") + if claims is not None: + return claims + return super().__getattribute__(name) + + return _Result() + + def _raise_import_error(content, event_type): raise ImportError(f"cannot import name 'redact' (formatting {content})") @@ -4685,6 +4717,179 @@ def formatter(content, event_type): assert payload_text not in json.dumps(row, default=str) assert payload_text not in caplog.text + @pytest.mark.parametrize( + "rethrow", + [True, False], + ids=["formatter_rethrows_it", "distinct_caller_exception"], + ) + async def test_failing_handler_never_prints_the_callers_active_exception( + self, + rethrow, + capsys, + monkeypatch, + mock_write_client, + invocation_context, + dummy_arrow_schema, + ): + """A failing handler cannot print an exception the caller is handling. + + ADK calls error callbacks while it handles the error, and a formatter can + re-raise that very exception. logging's handleError prints whatever is + being handled, so the warning is emitted while a constant stand-in is + handled instead. + """ + monkeypatch.setattr(logging, "raiseExceptions", True) + caller_failure = ValueError(f"the caller is handling {self.SECRET}") + + def formatter(content, event_type): + if rethrow: + raise caller_failure + raise ImportError("the formatter failed on its own") + + config = bigquery_agent_analytics_plugin.BigQueryLoggerConfig( + content_formatter=formatter + ) + closed_stream = io.StringIO() + closed_stream.close() + + with self._standard_handler_on_plugin_logger(closed_stream): + try: + raise caller_failure + except ValueError: + row, drop_stats = await self._log_user_message_contained( + config, mock_write_client, invocation_context, dummy_arrow_schema + ) + + stderr = capsys.readouterr().err + assert ( + row["content"] + == bigquery_agent_analytics_plugin._FORMATTER_FAILED_SENTINEL + ) + assert drop_stats.get("formatter_failed") == 1 + assert "--- Logging error ---" in stderr + assert self.SECRET not in stderr + + @pytest.mark.parametrize( + "hook_error", + [ + asyncio.CancelledError, + SystemExit, + KeyboardInterrupt, + RuntimeError, + None, + ], + ids=[ + "cancelled_error", + "system_exit", + "keyboard_interrupt", + "runtime_error", + "claims_to_be_a_dict", + ], + ) + @pytest.mark.parametrize("via", ["property", "getattribute"]) + async def test_result_that_lies_about_its_class_is_rejected_unrun( + self, + via, + hook_error, + mock_write_client, + invocation_context, + dummy_arrow_schema, + ): + """A result is judged by its real type, never by its own __class__. + + isinstance falls back to an object's __class__, which runs the object's + code: it can raise anything, or claim to be a dict. Its real type runs + none of that code, so such a result is simply an unsupported one, and + the formatter is not reported as having raised. + """ + tripwire = _Tripwire(hook_error or RuntimeError) + claims = dict if hook_error is None else None + + def formatter(content, event_type): + result = _result_with_class_hook(tripwire, via, claims) + tripwire.armed = hook_error is not None + return result + + config = bigquery_agent_analytics_plugin.BigQueryLoggerConfig( + content_formatter=formatter + ) + + def disarm(): + tripwire.armed = False + + row, drop_stats = await self._log_user_message_contained( + config, + mock_write_client, + invocation_context, + dummy_arrow_schema, + cleanup=disarm, + ) + + assert row["error_message"] == ( + "content_formatter returned unsupported type " + ) + assert ( + row["content"] + == bigquery_agent_analytics_plugin._FORMATTER_FAILED_SENTINEL + ) + assert drop_stats.get("formatter_failed") == 1 + + @pytest.mark.parametrize( + "payload_named", + [False, True], + ids=["async_formatter", "payload_named_coroutine"], + ) + async def test_rejected_coroutine_is_closed_without_a_warning( + self, + payload_named, + mock_write_client, + invocation_context, + dummy_arrow_schema, + ): + """A coroutine result is closed before it starts, so it never warns. + + Released unawaited, it would emit "coroutine '' was never + awaited", and a formatter can set that name from the content. + """ + ran = [] + + async def coroutine_body(): + ran.append("coroutine_body") + + if payload_named: + + def formatter(content, event_type): + coroutine = coroutine_body() + coroutine.__qualname__ = _identifier_from_content(content) + return coroutine + + else: + + async def formatter(content, event_type): + ran.append("formatter") + + config = bigquery_agent_analytics_plugin.BigQueryLoggerConfig( + content_formatter=formatter + ) + + with warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always") + row, drop_stats = await self._log_user_message_contained( + config, mock_write_client, invocation_context, dummy_arrow_schema + ) + gc.collect() + + assert row["error_message"] == ( + "content_formatter returned unsupported type coroutine" + ) + assert drop_stats.get("formatter_failed") == 1 + assert not ran + messages = [str(warning.message) for warning in caught] + assert not [message for message in messages if "never awaited" in message] + assert not [ + message for message in messages if self.PAYLOAD_IDENTIFIER in message + ] + @pytest.mark.parametrize( "injected", [ @@ -4712,6 +4917,8 @@ def formatter(content, event_type): pytest.param("log_handler", False, id="log_handler-returned"), pytest.param("log_filter", True, id="log_filter-raised"), pytest.param("log_filter", False, id="log_filter-returned"), + pytest.param("admit", False, id="admit-returned"), + pytest.param("close", False, id="close-returned"), ], ) async def test_a_raise_anywhere_in_diagnosis_leaves_the_row_intact( @@ -4724,13 +4931,14 @@ async def test_a_raise_anywhere_in_diagnosis_leaves_the_row_intact( dummy_arrow_schema, caplog, ): - """Whatever any diagnostic step raises, the row is written unchanged. + """Whatever any step after the formatter call raises, the row survives. - Diagnosis (naming the failed class, rendering the debug traceback, and - emitting the warning through the logger's filters and handlers) runs - after the failure is handled, behind one boundary. Each step here raises - each kind of exception, and the sentinel row, its drop count, and a - payload-free error_message must still come out, with nothing escaping. + Judging the result, closing a rejected generator, naming the failed + class, rendering the debug traceback, and emitting the warning through + the logger's filters and handlers all run behind one boundary. Each step + here raises each kind of exception, and the sentinel row, its drop count, + and a payload-free error_message must still come out, with nothing + escaping. """ plugin_module = bigquery_agent_analytics_plugin plugin_logger = logging.getLogger("google_adk." + plugin_module.__name__) @@ -4748,9 +4956,25 @@ class _RaisingHandler(logging.Handler): def emit(self, record): raise_for_formatter_warnings(record) + def return_started_generator(content, event_type): + def generator(): + try: + yield "started" + finally: + raise_injected() + + started = generator() + next(started) + return started + + if raised: + formatter = _raise_import_error + elif step == "close": + formatter = return_started_generator + else: + formatter = _return_tuple config = plugin_module.BigQueryLoggerConfig( - content_formatter=_raise_import_error if raised else _return_tuple, - debug_content_formatter_errors=True, + content_formatter=formatter, debug_content_formatter_errors=True ) handler = _RaisingHandler() injections = { @@ -4762,6 +4986,10 @@ def emit(self, record): ), "log_handler": contextlib.nullcontext(), "log_filter": contextlib.nullcontext(), + "admit": mock.patch.object( + plugin_module, "_natively_parsed", raise_injected, create=True + ), + "close": contextlib.nullcontext(), } if step == "log_handler": plugin_logger.addHandler(handler) @@ -4777,7 +5005,7 @@ def emit(self, record): plugin_logger.removeFilter(raise_for_formatter_warnings) outcome = "raised" if raised else "returned unsupported type" - if step == "label": + if step in ("label", "admit", "close"): expected_error_message = f"content_formatter {outcome} " elif raised: expected_error_message = "content_formatter raised ImportError" From f64500aa648ede9b2009c60d1b1e773eb96785d2 Mon Sep 17 00:00:00 2001 From: Haiyuan Cao Date: Tue, 29 Sep 2026 09:07:18 -0700 Subject: [PATCH 06/29] fix(plugins): honor genuine interrupts and isolate all plugin logging Review of the previous commit found two remaining gaps, and two guards that no test pinned: - A KeyboardInterrupt or SystemExit delivered while a content_formatter failure was being described was absorbed. A real SIGINT, or a SIGTERM whose handler calls sys.exit while a log handler blocks on I/O, left the row written but the process running, unaware of the signal. - Only the formatter-failure warning ran while a stand-in exception was handled. Any other plugin warning, such as the parser's, could still have a failing handler print the exception the caller is handling, as ADK is when it runs error callbacks. Interrupts are now sorted by the code that raised them, because Python cannot tell a signal from a direct raise. The failed class's and the rejected result's own code runs only while the debug traceback is rendered and while a rejected generator is closed. Each of those steps has its own guard, and whatever it raises, interrupts included, is contained, so the content under redaction cannot end the agent run. Everywhere else in the boundary only plugin code and the application's log filters and handlers run, so a KeyboardInterrupt or SystemExit there came from a signal or from the application. _settle_formatter_outcome returns it as a new exception without text (a SystemExit keeps an int exit code), and _log_event raises it once the row has been handed to the writer. CancelledError stays contained: nothing in the boundary awaits, so it cannot be a real cancellation. The trade-off is that a signal landing while payload-controlled code runs is still absorbed. The module logger's handle() now runs every record's filters and handlers while a constant stand-in exception is handled, with its __context__ cleared. So no plugin log call can print the caller's exception through handleError, or through a handler that walks __context__ and ignores __suppress_context__. Records are unchanged: Logger._log resolves exc_info and the calling function before handle() runs. This replaces the stand-in around the formatter warning alone. The boundary's docstring now bounds its never-drop claim to failed or rejected results. An admitted str, dict, or list subclass goes on to the parser, whose boundary catches Exception only. Tests: - the injection property test expects an interrupt from any step but closing to be raised after the row, as a new exception, and checks that a rejected generator's cleanup never reaches the unraisable hook; - a real SIGINT and SIGTERM delivered while the warning is emitted; - a parse-failure warning with a closed handler, and with a handler that walks __context__, while the caller handles a payload exception; - the plugin's error logs keep their own exc_info; - supported results (None, a str subclass whose hooks raise, dict and list subclasses, and exact Content, Part, and LlmRequest) pass through unchanged. Removing any new guard fails a test. Refs: https://github.com/GoogleCloudPlatform/BigQuery-Agent-Analytics-SDK/issues/485 (item A) Co-Authored-By: Claude Opus 5.5 (1M context) --- .../bigquery_agent_analytics_plugin.py | 248 +++++++++---- .../test_bigquery_agent_analytics_plugin.py | 340 +++++++++++++++++- 2 files changed, 511 insertions(+), 77 deletions(-) diff --git a/src/google/adk/plugins/bigquery_agent_analytics_plugin.py b/src/google/adk/plugins/bigquery_agent_analytics_plugin.py index 373af9abe98..f64786aaefe 100644 --- a/src/google/adk/plugins/bigquery_agent_analytics_plugin.py +++ b/src/google/adk/plugins/bigquery_agent_analytics_plugin.py @@ -112,7 +112,43 @@ from ..agents.invocation_context import InvocationContext from ..events.event import Event + +class _LoggingStandIn(Exception): + """The exception being handled while this module's log records are handled. + + logging's ``handleError``, and handlers that report the current exception, + print the exception being handled. Without this stand-in that can be one + the caller is handling, such as the error ADK passes to an error callback, + and its text can carry the content this plugin keeps out of logs. + """ + + +def _handle_records_with_a_stand_in(target: logging.Logger) -> None: + """Makes ``target`` run its filters and handlers while handling a stand-in. + + The stand-in's ``__context__`` is cleared, so no exception chain printed + while a record is handled can reach the caller's exception, even by a + handler that ignores ``__suppress_context__``. Records are unchanged: + ``Logger._log`` resolves the calling function and any ``exc_info`` before + it calls ``handle``. + + Args: + target: The logger whose records to handle this way. + """ + handle = type(target).handle + + def handle_with_stand_in(record: logging.LogRecord) -> None: + try: + raise _LoggingStandIn + except _LoggingStandIn as stand_in: + stand_in.__context__ = None + handle(target, record) + + target.handle = handle_with_stand_in # type: ignore[method-assign] + + logger: logging.Logger = logging.getLogger("google_adk." + __name__) +_handle_records_with_a_stand_in(logger) # Bumped when the schema changes (1 → 2 → 3 …). Used as a table # label for governance and to decide whether auto-upgrade should run. @@ -1120,13 +1156,27 @@ def _render_formatter_traceback(error: BaseException) -> str: return "[traceback could not be rendered]" -class _WarningIsolation(Exception): - """The exception being handled while the formatter-failure warning is logged. +# SystemExit's own ``code`` descriptor; reading it through the descriptor +# runs no code of a SystemExit subclass. +_SYSTEM_EXIT_CODE = SystemExit.__dict__["code"] - A log handler that fails calls ``handleError``, which prints the exception - being handled. Without this stand-in that could be the formatter's - exception, or one the caller is handling, and either can carry the content. + +def _fresh_interrupt(interrupt: BaseException) -> BaseException: + """Returns a new KeyboardInterrupt or SystemExit that carries no text. + + A SystemExit keeps its exit code only when the code is an int or None; + any other code, such as a message, becomes 1, Python's failure status. + + Args: + interrupt: A KeyboardInterrupt or SystemExit, possibly a subclass. + + Returns: + The fresh exception to raise in its place. """ + if issubclass(type(interrupt), KeyboardInterrupt): + return KeyboardInterrupt() + code = _SYSTEM_EXIT_CODE.__get__(interrupt) + return SystemExit(code if code is None or type(code) is int else 1) def _natively_parsed(formatted: Any, result_type: type) -> bool: @@ -1162,10 +1212,10 @@ def _settle_formatter_outcome( *, event_type: str, debug: bool, -) -> tuple[Any, Optional[str]]: +) -> tuple[Any, Optional[str], Optional[BaseException]]: """Decides what the row logs after the content_formatter call, fail closed. - Everything after the call happens behind this one boundary: + These steps happen behind this one boundary: - judging a returned result by its real type; - closing a rejected coroutine or generator; - naming the failed class; @@ -1173,22 +1223,25 @@ def _settle_formatter_outcome( - emitting the warning through whatever filters and handlers are configured. - Whatever any step raises, of any type, is contained here. The outcome then - falls back to the sentinel and a constant note, so nothing after the call - can drop the row, leave the sentinel out, or escape the callback. - - BaseException is contained on purpose. These steps run code that the - failed class or the rejected result controls, and code in filters and - handlers. A KeyboardInterrupt or SystemExit raised there cannot be told - apart from one that a signal handler delivered, and letting it through - would let the content under redaction abort the agent run and lose the - row. The cost is that a signal delivered while these steps run is absorbed - and the process does not act on it. That includes a signal that arrives - while a configured handler blocks, for example on network I/O. An - orchestrator that follows SIGTERM with SIGKILL then loses the rows still - buffered. Nothing here awaits, so a real asyncio cancellation is never - absorbed, and interrupts raised by the formatter call itself still - propagate. + Whatever those steps raise is contained here, and the outcome falls back + to the sentinel and a constant note, so no failure while describing a + failed or rejected result can drop the row or leave the sentinel out. A + result the parser logs natively leaves unchanged, and the parser's own + boundary, which catches Exception only, applies to it. + + Interrupts are sorted by the code that raised them, because Python cannot + tell a KeyboardInterrupt or SystemExit that a signal handler delivered + from one raised directly: code can even signal its own process. Code that + the failed class or the rejected result controls runs only while the + traceback is rendered and while a rejected generator is closed, and + anything raised there, interrupts included, is contained, so that the + content under redaction cannot end the agent run. A signal that lands + there is absorbed. Everywhere else only this module's code and the + application's log filters and handlers run. There a KeyboardInterrupt or + SystemExit came from a signal or from the application, so it is returned, + as a fresh exception with no text, for the caller to raise once the row is + written. CancelledError is always contained: nothing here awaits, so it + cannot be a real cancellation. Args: formatted: What the formatter returned; ignored when ``failure`` is set. @@ -1197,10 +1250,11 @@ def _settle_formatter_outcome( debug: Whether to append the rendered traceback to the warning. Returns: - ``(formatted, None)`` when the parser can log the result as it is. A str - subclass is normalized to the exact built-in. Otherwise - ``(_FORMATTER_FAILED_SENTINEL, note)``, where ``note`` is the text for - the error_message column. + ``(formatted, None, None)`` when the parser can log the result as it + is; a str subclass is normalized to the exact built-in. Otherwise + ``(_FORMATTER_FAILED_SENTINEL, note, interrupt)``, where ``note`` is the + text for the error_message column and ``interrupt`` is None or the + interrupt to raise after the row is written. """ outcome = "raised" if failure is not None else "returned unsupported type" note = f"content_formatter {outcome} " @@ -1212,46 +1266,47 @@ def _settle_formatter_outcome( if _natively_parsed(formatted, failed_type): if failed_type is not str and issubclass(failed_type, str): formatted = str.__str__(formatted) - return formatted, None + return formatted, None, None # A non-native result would reach the parser's str() fallback, where # a payload-controlled __str__ can republish the content, so it is # rejected. Close a coroutine first: released unstarted, it warns # "coroutine '' was never awaited", and the formatter can set # that name from the content. A generator is closed too, so that its # cleanup code runs here, contained, rather than at collection. - if failed_type is CoroutineType: - CoroutineType.close(formatted) - elif failed_type is GeneratorType: - GeneratorType.close(formatted) + try: + if failed_type is CoroutineType: + CoroutineType.close(formatted) + elif failed_type is GeneratorType: + GeneratorType.close(formatted) + except BaseException: + # Closing ran the result's own code; see the docstring. + pass note = _formatter_failure_message(failed_type, raised=failure is not None) - try: - raise _WarningIsolation from None - except _WarningIsolation: - if failure is None: - logger.warning( - "Content formatter returned an unsupported result type for" - " event %s; writing sentinel instead of original content.", - event_type, - ) - elif debug: - logger.warning( - "Content formatter failed for event %s; writing sentinel" - " instead of original content. Debug traceback:\n%s", - event_type, - _render_formatter_traceback(failure), - ) - else: - logger.warning( - "Content formatter failed for event %s; writing sentinel" - " instead of original content.", - event_type, - ) - except BaseException: - # Contained whatever it is, for the reasons in the docstring. It is not - # reported either: the logger may be what failed, and the exception can - # carry the content. - pass - return _FORMATTER_FAILED_SENTINEL, note + if failure is None: + logger.warning( + "Content formatter returned an unsupported result type for" + " event %s; writing sentinel instead of original content.", + event_type, + ) + elif debug: + logger.warning( + "Content formatter failed for event %s; writing sentinel" + " instead of original content. Debug traceback:\n%s", + event_type, + _render_formatter_traceback(failure), + ) + else: + logger.warning( + "Content formatter failed for event %s; writing sentinel" + " instead of original content.", + event_type, + ) + except BaseException as error: + # Nothing is reported: the logger may be what failed, and the exception + # can carry the content. See the docstring for which interrupts return. + if issubclass(type(error), (KeyboardInterrupt, SystemExit)): + return _FORMATTER_FAILED_SENTINEL, note, _fresh_interrupt(error) + return _FORMATTER_FAILED_SENTINEL, note, None # Recursion bound for _recursive_smart_truncate: id()-based cycle detection @@ -2357,12 +2412,17 @@ class BigQueryLoggerConfig: are named. Any other class, including one your own code defines, is described by its nearest named ancestor, for example ``content_formatter raised ``, and a class - that cannot be read at all as ````. Judging the - result, describing the failure, and logging its warning are best - effort: whatever they raise is contained, so they never drop the - row or change the sentinel or the counter. A ``KeyboardInterrupt``, - ``SystemExit``, or ``asyncio.CancelledError`` raised by the - formatter call itself still propagates, and no row is written. An + that cannot be read at all as ````. Describing a + failed or rejected result and logging its warning are best effort: + nothing they raise drops the row or changes the sentinel or the + counter. A ``KeyboardInterrupt`` or ``SystemExit`` that a signal + handler or your own log handlers and filters raise meanwhile is + raised again once the row is written, as a new exception without + text; one raised by the failed class's or the rejected result's own + code is contained, as is a signal that arrives while that code runs. + A ``KeyboardInterrupt``, ``SystemExit``, or + ``asyncio.CancelledError`` raised by the formatter call itself still + propagates, and no row is written. An event that already carries an ``error_message``, such as a ``TOOL_ERROR``, keeps it first, followed by ``; `` and the formatter failure. @@ -7411,6 +7471,49 @@ async def _log_event( is_truncated: Whether the content is already truncated. event_data: Typed container for structured fields and extra attributes. Defaults to ``EventData()`` when not provided. + + Raises: + KeyboardInterrupt: A signal handler, or a log handler or filter, + raised one while a content_formatter failure was being described. + A new one without text is raised after the row was handed to the + writer; see ``_settle_formatter_outcome``. + SystemExit: Likewise; it keeps the exit code only if that is an int. + """ + interrupts: list[BaseException] = [] + try: + await self._log_event_row( + event_type, + callback_context, + raw_content, + is_truncated, + event_data, + interrupts, + ) + finally: + if interrupts: + # Raised only now, after the row was handed to the writer, so that + # neither the row nor the signal is lost. + raise interrupts[0] from None + + async def _log_event_row( + self, + event_type: str, + callback_context: CallbackContext, + raw_content: Any, + is_truncated: bool, + event_data: Optional[EventData], + interrupts: list[BaseException], + ) -> None: + """Builds the row for ``_log_event`` and hands it to the writer. + + Args: + event_type: As for ``_log_event``. + callback_context: As for ``_log_event``. + raw_content: As for ``_log_event``. + is_truncated: As for ``_log_event``. + event_data: As for ``_log_event``. + interrupts: Receives an interrupt deferred while a content_formatter + failure was described, for ``_log_event`` to raise afterwards. """ if not self.config.enabled or self._is_shutting_down: return @@ -7491,11 +7594,12 @@ async def _log_event( # debug_content_formatter_errors opts in. failure = e # Everything after the call runs behind _settle_formatter_outcome's - # one boundary, which contains whatever it raises: judging the result, - # closing a rejected coroutine, naming the class, and logging. It runs - # after the except block, so the formatter's exception is no longer - # the one being handled. - raw_content, formatter_error = _settle_formatter_outcome( + # one boundary: judging the result, closing a rejected coroutine, + # naming the class, and logging. Nothing it raises reaches here; an + # interrupt it sets aside comes back to be raised once the row is + # written. It runs after the except block, so the formatter's + # exception is no longer the one being handled. + raw_content, formatter_error, interrupt = _settle_formatter_outcome( formatted, failure, event_type=event_type, @@ -7503,6 +7607,8 @@ async def _log_event( ) if formatter_error is not None: self._count_local_drop("formatter_failed") + if interrupt is not None: + interrupts.append(interrupt) # The except clause would have dropped this reference itself: the # exception's traceback holds this frame, which holds the exception. formatted = failure = None diff --git a/tests/unittests/plugins/test_bigquery_agent_analytics_plugin.py b/tests/unittests/plugins/test_bigquery_agent_analytics_plugin.py index 36dff868632..96a18180489 100644 --- a/tests/unittests/plugins/test_bigquery_agent_analytics_plugin.py +++ b/tests/unittests/plugins/test_bigquery_agent_analytics_plugin.py @@ -23,6 +23,7 @@ import logging import os import pickle +import signal import sys import threading import time @@ -3982,6 +3983,53 @@ def __getattribute__(self, name): return _Result() +def _hostile_str(tripwire, text): + """Returns a str subclass instance whose own hooks fire once armed.""" + + class _HostileStr(str): + + @property + def __class__(self): + tripwire.fire("str __class__ property") + return type(self) + + def __getattribute__(self, name): + tripwire.fire(f"str __getattribute__({name!r})") + return super().__getattribute__(name) + + def __str__(self): + tripwire.fire("str __str__") + return super().__str__() + + def __len__(self): + tripwire.fire("str __len__") + return super().__len__() + + def __hash__(self): + tripwire.fire("str __hash__") + return super().__hash__() + + def __format__(self, spec): + tripwire.fire("str __format__") + return super().__format__(spec) + + return _HostileStr(text) + + +class _ContextWalkingHandler(logging.Handler): + """Writes the handled exception's chain, ignoring __suppress_context__.""" + + def __init__(self, stream): + super().__init__() + self.stream = stream + + def emit(self, record): + handled = sys.exc_info()[1] + while handled is not None: + self.stream.write(f"{type(handled).__name__}: {handled}\n") + handled = handled.__context__ + + def _raise_import_error(content, event_type): raise ImportError(f"cannot import name 'redact' (formatting {content})") @@ -4155,6 +4203,45 @@ async def _log_user_message_contained( f"{type(escaped).__name__} escaped the plugin callback", pytrace=False ) + async def _log_user_message_catching( + self, + config, + mock_write_client, + invocation_context, + dummy_arrow_schema, + *, + cleanup=None, + ): + """Logs SECRET as a user message; returns the row, stats, and escapee. + + Whatever the callback raises is caught and returned, and the row is still + flushed and read afterwards, so a test can check that an interrupt was + raised only after the row reached the writer. cleanup runs before + anything reads the escaped exception. + """ + escaped = None + async with managed_plugin( + PROJECT_ID, DATASET_ID, table_id=TABLE_ID, config=config + ) as plugin: + await plugin._ensure_started() + mock_write_client.append_rows.reset_mock() + bigquery_agent_analytics_plugin.TraceManager.push_span(invocation_context) + try: + await plugin.on_user_message_callback( + invocation_context=invocation_context, + user_message=types.Content(parts=[types.Part(text=self.SECRET)]), + ) + except BaseException as error: # pylint: disable=broad-exception-caught + escaped = error + finally: + if cleanup is not None: + cleanup() + await plugin.flush() + row = await _get_captured_event_dict_async( + mock_write_client, dummy_arrow_schema + ) + return row, plugin.get_drop_stats(), escaped + @staticmethod def _formatter_warnings(caplog): return [ @@ -4717,6 +4804,219 @@ def formatter(content, event_type): assert payload_text not in json.dumps(row, default=str) assert payload_text not in caplog.text + @pytest.mark.parametrize( + "result", + [ + "none", + "hostile_str_subclass", + "dict_subclass", + "list_subclass", + "content", + "part", + "llm_request", + ], + ) + async def test_supported_results_pass_through_unchanged( + self, result, mock_write_client, invocation_context, dummy_arrow_schema + ): + """Results the parser logs natively are logged, and nothing is counted. + + A str subclass is copied to the exact built-in first, so the parser never + runs the subclass's own hooks. + """ + tripwire = _Tripwire(asyncio.CancelledError) + + class _Dict(dict): + pass + + class _List(list): + pass + + def formatter(content, event_type): + if result == "hostile_str_subclass": + value = _hostile_str(tripwire, "redacted text") + tripwire.armed = True + return value + return { + "none": None, + "dict_subclass": _Dict(redacted="text"), + "list_subclass": _List(["redacted"]), + "content": types.Content(parts=[types.Part(text="redacted")]), + "part": types.Part(text="redacted"), + "llm_request": llm_request_lib.LlmRequest(), + }[result] + + config = bigquery_agent_analytics_plugin.BigQueryLoggerConfig( + content_formatter=formatter + ) + + def disarm(): + tripwire.armed = False + + row, drop_stats = await self._log_user_message_contained( + config, + mock_write_client, + invocation_context, + dummy_arrow_schema, + cleanup=disarm, + ) + + assert drop_stats.get("formatter_failed", 0) == 0 + assert row["error_message"] is None + assert ( + row["content"] + != bigquery_agent_analytics_plugin._FORMATTER_FAILED_SENTINEL + ) + if result == "hostile_str_subclass": + assert row["content"] == "redacted text" + + @pytest.mark.skipif(sys.platform == "win32", reason="needs POSIX signals") + @pytest.mark.parametrize( + ("signal_name", "expected"), + [("SIGTERM", SystemExit), ("SIGINT", KeyboardInterrupt)], + ) + async def test_genuine_signal_during_the_warning_is_raised_after_the_row( + self, + signal_name, + expected, + mock_write_client, + invocation_context, + dummy_arrow_schema, + ): + """A real signal that lands while the warning is emitted is not lost. + + The row is handed to the writer first; the interrupt is raised after it. + The SIGTERM handler exits with a code, as orchestrated shutdowns do. + """ + signum = getattr(signal, signal_name) + + def exit_on_sigterm(signum, frame): + sys.exit(143) + + handlers = { + "SIGTERM": exit_on_sigterm, + "SIGINT": signal.default_int_handler, + } + plugin_logger = logging.getLogger( + "google_adk." + bigquery_agent_analytics_plugin.__name__ + ) + + class _SlowHandler(logging.Handler): + + def emit(self, record): + if record.getMessage().startswith("Content formatter "): + threading.Timer(0.05, os.kill, (os.getpid(), signum)).start() + time.sleep(5) + + config = bigquery_agent_analytics_plugin.BigQueryLoggerConfig( + content_formatter=_raise_import_error + ) + slow_handler = _SlowHandler() + previous_handler = signal.signal(signum, handlers[signal_name]) + plugin_logger.addHandler(slow_handler) + try: + row, drop_stats, escaped = await self._log_user_message_catching( + config, mock_write_client, invocation_context, dummy_arrow_schema + ) + finally: + plugin_logger.removeHandler(slow_handler) + signal.signal(signum, previous_handler) + + assert type(escaped) is expected + if expected is SystemExit: + assert escaped.code == 143 + assert row["error_message"] == "content_formatter raised ImportError" + assert ( + row["content"] + == bigquery_agent_analytics_plugin._FORMATTER_FAILED_SENTINEL + ) + assert drop_stats.get("formatter_failed") == 1 + + @pytest.mark.parametrize("handler_kind", ["closed_stream", "context_walker"]) + async def test_plugin_warnings_never_print_the_callers_exception( + self, + handler_kind, + capsys, + monkeypatch, + mock_write_client, + invocation_context, + dummy_arrow_schema, + ): + """No plugin warning can print an exception the caller is handling. + + Here the content parser fails, an ordinary plugin warning unrelated to + content_formatter, while the caller handles an error, as ADK does when + it runs error callbacks. + """ + monkeypatch.setattr(logging, "raiseExceptions", True) + config = bigquery_agent_analytics_plugin.BigQueryLoggerConfig() + walked = io.StringIO() + if handler_kind == "closed_stream": + stream = io.StringIO() + stream.close() + handler = logging.StreamHandler(stream) + else: + handler = _ContextWalkingHandler(walked) + plugin_logger = logging.getLogger( + "google_adk." + bigquery_agent_analytics_plugin.__name__ + ) + caller_failure = ValueError(f"the caller is handling {self.SECRET}") + + plugin_logger.addHandler(handler) + try: + with mock.patch.object( + bigquery_agent_analytics_plugin.HybridContentParser, + "parse", + side_effect=RuntimeError("the parser failed"), + ): + try: + raise caller_failure + except ValueError: + row, _ = await self._log_user_message_contained( + config, mock_write_client, invocation_context, dummy_arrow_schema + ) + finally: + plugin_logger.removeHandler(handler) + + stderr = capsys.readouterr().err + printed = stderr + walked.getvalue() + assert row["content"] == "[CONTENT_PARSE_FAILED]" + if handler_kind == "closed_stream": + assert "--- Logging error ---" in stderr + else: + assert walked.getvalue(), "the handler ran with no exception handled" + assert self.SECRET not in printed + + async def test_plugin_error_logs_keep_their_own_exception( + self, invocation_context, caplog + ): + """Handling a stand-in never replaces the exception a log records.""" + plugin = bigquery_agent_analytics_plugin.BigQueryAgentAnalyticsPlugin( + PROJECT_ID, DATASET_ID, table_id=TABLE_ID + ) + + try: + with ( + mock.patch.object( + plugin, "_log_event", side_effect=RuntimeError("write failed") + ), + caplog.at_level(logging.ERROR), + ): + await plugin.on_user_message_callback( + invocation_context=invocation_context, + user_message=types.Content(parts=[types.Part(text="hello")]), + ) + finally: + await plugin.shutdown() + + records = [ + record + for record in caplog.records + if "plugin error in on_user_message_callback" in record.getMessage() + ] + assert len(records) == 1 + assert records[0].exc_info[0] is RuntimeError + @pytest.mark.parametrize( "rethrow", [True, False], @@ -4937,8 +5237,12 @@ async def test_a_raise_anywhere_in_diagnosis_leaves_the_row_intact( class, rendering the debug traceback, and emitting the warning through the logger's filters and handlers all run behind one boundary. Each step here raises each kind of exception, and the sentinel row, its drop count, - and a payload-free error_message must still come out, with nothing - escaping. + and a payload-free error_message must still come out. + + Closing runs the rejected generator's own code, so whatever it raises is + contained. Every other step runs only plugin or application code, so a + KeyboardInterrupt or SystemExit there is honored: it is raised after the + row is written, as a fresh exception carrying no text. """ plugin_module = bigquery_agent_analytics_plugin plugin_logger = logging.getLogger("google_adk." + plugin_module.__name__) @@ -4991,22 +5295,46 @@ def generator(): ), "close": contextlib.nullcontext(), } + unraisable = [] + + def record_unraisable(hook_args): + unraisable.append(hook_args.exc_value) + if step == "log_handler": plugin_logger.addHandler(handler) if step == "log_filter": plugin_logger.addFilter(raise_for_formatter_warnings) try: with injections[step], caplog.at_level(logging.WARNING): - row, drop_stats = await self._log_user_message_contained( - config, mock_write_client, invocation_context, dummy_arrow_schema - ) + with mock.patch.object(sys, "unraisablehook", record_unraisable): + row, drop_stats, escaped = await self._log_user_message_catching( + config, mock_write_client, invocation_context, dummy_arrow_schema + ) + gc.collect() finally: plugin_logger.removeHandler(handler) plugin_logger.removeFilter(raise_for_formatter_warnings) + # The rejected generator is closed inside the boundary. Released unclosed, + # its cleanup would raise at the unraisable hook, which prints the error. + assert not [ + error + for error in unraisable + if type(error) is injected and error.args == (f"injected into {step}",) + ] + if step != "close" and injected in (KeyboardInterrupt, SystemExit): + assert type(escaped) is injected + # Fresh, with no text: an exit code survives only if it is an int. + assert escaped.args == (() if injected is KeyboardInterrupt else (1,)) + else: + assert escaped is None outcome = "raised" if raised else "returned unsupported type" - if step in ("label", "admit", "close"): + if step in ("label", "admit"): expected_error_message = f"content_formatter {outcome} " + elif step == "close": + expected_error_message = ( + "content_formatter returned unsupported type generator" + ) elif raised: expected_error_message = "content_formatter raised ImportError" else: From 9f16d70b9df2da5b4ebc375f2a2953713c14b2b0 Mon Sep 17 00:00:00 2001 From: Haiyuan Cao Date: Tue, 29 Sep 2026 09:11:41 -0700 Subject: [PATCH 07/29] docs(plugins): note that closing a rejected coroutine is contained too The boundary's docstring said the rejected result's own code runs only while the debug traceback is rendered and while a rejected generator is closed. The same guard also covers closing a rejected coroutine, which runs the coroutine's own cleanup if it was started. Wording only. Refs: https://github.com/GoogleCloudPlatform/BigQuery-Agent-Analytics-SDK/issues/485 (item A) Co-Authored-By: Claude Opus 5.5 (1M context) --- src/google/adk/plugins/bigquery_agent_analytics_plugin.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/src/google/adk/plugins/bigquery_agent_analytics_plugin.py b/src/google/adk/plugins/bigquery_agent_analytics_plugin.py index f64786aaefe..bd9fd4b62aa 100644 --- a/src/google/adk/plugins/bigquery_agent_analytics_plugin.py +++ b/src/google/adk/plugins/bigquery_agent_analytics_plugin.py @@ -1233,7 +1233,8 @@ def _settle_formatter_outcome( tell a KeyboardInterrupt or SystemExit that a signal handler delivered from one raised directly: code can even signal its own process. Code that the failed class or the rejected result controls runs only while the - traceback is rendered and while a rejected generator is closed, and + traceback is rendered and while a rejected coroutine or generator is + closed, and anything raised there, interrupts included, is contained, so that the content under redaction cannot end the agent run. A signal that lands there is absorbed. Everywhere else only this module's code and the From c50ade39c88f391032887bcd1ba5e9299697d8c8 Mon Sep 17 00:00:00 2001 From: George Weale Date: Tue, 29 Sep 2026 10:05:12 -0700 Subject: [PATCH 08/29] feat: honor tool_thread_pool_config for sync tools outside live mode A synchronous function tool blocks the event loop while it runs, and RunConfig.tool_thread_pool_config, the setting that moves tools off the loop, only applied in live mode. With it set, run_async now runs a synchronous function tool's function on the tool thread pool when an LlmAgent calls the tool, while async tools and tools used directly as Workflow nodes stay on the event loop, and without it nothing changes. Co-authored-by: George Weale PiperOrigin-RevId: 990376105 --- src/google/adk/agents/run_config.py | 32 ++-- .../adk/flows/llm_flows/tools/_caller.py | 83 +++++++--- .../llm_flows/tools/test_functions_simple.py | 148 ++++++++++++++++++ tests/unittests/testing_utils.py | 13 +- 4 files changed, 241 insertions(+), 35 deletions(-) diff --git a/src/google/adk/agents/run_config.py b/src/google/adk/agents/run_config.py index e6910c1a3ce..de360fc72ad 100644 --- a/src/google/adk/agents/run_config.py +++ b/src/google/adk/agents/run_config.py @@ -55,6 +55,9 @@ def _default_max_llm_calls() -> int: class ToolThreadPoolConfig(BaseModel): """Configuration for the tool thread pool executor. + Set through `RunConfig.tool_thread_pool_config`. Outside live mode only + synchronous function tools an `LlmAgent` calls use the pool. + Attributes: max_workers: Maximum number of worker threads in the pool. Defaults to 4. """ @@ -186,21 +189,32 @@ class RunConfig(BaseModel): """Saves live video and audio data to session and artifact service.""" tool_thread_pool_config: Optional[ToolThreadPoolConfig] = None - """Configuration for running tools in a thread pool for live mode. + """Configuration for running tools in a thread pool. When set, tool executions will run in a separate thread pool executor - instead of the main event loop. When None (default), tools run in the - main event loop. One pool serves every invocation running on the same event - loop and is shut down once that loop is gone, so its worker threads do not - outlive it. + instead of the main event loop, for the tools described below. When None + (default), tools run in the main event loop. One pool serves every + invocation running on the same event loop and is shut down once that loop + is gone, so its worker threads do not outlive it. - This helps keep the event loop responsive for: + In live mode, this helps keep the event loop responsive for: - User interruptions to be processed immediately - Model responses to continue being received - Both sync and async tools are supported. Async tools are run in a new event - loop within the background thread, which helps catch blocking I/O mistakenly - used inside async functions. + In live mode, both sync and async tools are supported. Async tools are run + in a new event loop within the background thread, which helps catch blocking + I/O mistakenly used inside async functions. + + Outside live mode, only synchronous function tools an `LlmAgent` calls use + the pool, and only their own synchronous callables run there: the tool + function and a callable `require_confirmation`. Argument handling and + callbacks stay on the event loop, and async tools run on the event loop as + they do without this config. A tool used directly as a `Workflow` node also + runs on the event loop, because a tool node calls the tool itself rather + than through the `LlmAgent` tool pipeline this config applies to. Parallel + calls to synchronous function tools then overlap, so their functions must + be safe to run at the same time and on a thread other than the event + loop's. IMPORTANT - GIL (Global Interpreter Lock) Considerations: diff --git a/src/google/adk/flows/llm_flows/tools/_caller.py b/src/google/adk/flows/llm_flows/tools/_caller.py index 54e39780fc1..c99e1a07bee 100644 --- a/src/google/adk/flows/llm_flows/tools/_caller.py +++ b/src/google/adk/flows/llm_flows/tools/_caller.py @@ -20,7 +20,9 @@ import base64 import binascii from collections.abc import Awaitable +from collections.abc import Iterator from concurrent.futures import ThreadPoolExecutor +import contextlib import contextvars import copy import dataclasses @@ -180,6 +182,34 @@ def _is_sync_tool(tool: BaseTool) -> bool: ) +@contextlib.contextmanager +def _use_executor_for_sync_callables( + executor: ThreadPoolExecutor, +) -> Iterator[None]: + """Binds a sync callable runner that calls each callable on ``executor``. + + The callable runs with a copy of the caller's context variables and with no + runner bound, so a nested call it makes runs inline on its worker thread. + """ + + async def run_sync_callable( + target: Callable[..., Any], call_args: dict[str, Any] + ) -> Any: + call_context = contextvars.copy_context() + + def invoke() -> Any: + with _use_sync_callable_runner(None): + return target(**call_args) + + return await asyncio.get_running_loop().run_in_executor( + executor, + lambda: call_context.run(invoke), + ) + + with _use_sync_callable_runner(run_sync_callable): + yield + + async def _call_tool_in_thread_pool( tool: BaseTool, args: dict[str, Any], @@ -209,22 +239,7 @@ async def _call_tool_in_thread_pool( executor = _get_tool_thread_pool(max_workers) if _is_sync_tool(tool) and isinstance(tool, FunctionTool): - - async def run_sync_callable( - target: Callable[..., Any], call_args: dict[str, Any] - ) -> Any: - call_context = contextvars.copy_context() - - def invoke() -> Any: - with _use_sync_callable_runner(None): - return target(**call_args) - - return await loop.run_in_executor( - executor, - lambda: call_context.run(invoke), - ) - - with _use_sync_callable_runner(run_sync_callable): + with _use_executor_for_sync_callables(executor): return await tool.run_async(args=args, tool_context=tool_context) ctx = contextvars.copy_context() @@ -934,16 +949,38 @@ async def _execute_single_prepared_call_async( the tool unless one of them answered the call, run the after-tool callbacks, and turn the result into an event. State modifications stay thread safe because each call owns its own ToolContext. + + With `RunConfig.tool_thread_pool_config` set, a synchronous `FunctionTool` + calls its function on the tool thread pool. Every other tool, async function + tools included, runs on the event loop as it does without the config. """ - return await _execute_single_prepared_call( - invocation_context, - prepared_call, - agent, - tool_runner=lambda: _call_tool_async( - prepared_call.tool, + tool = prepared_call.tool + run_config = invocation_context.run_config + thread_pool_config = ( + run_config.tool_thread_pool_config if run_config else None + ) + + async def call_tool() -> object: + sync_callables: contextlib.AbstractContextManager[None] = ( + contextlib.nullcontext() + ) + if ( + thread_pool_config is not None + and _is_sync_tool(tool) + and isinstance(tool, FunctionTool) + ): + sync_callables = _use_executor_for_sync_callables( + _get_tool_thread_pool(thread_pool_config.max_workers) + ) + with sync_callables: + return await _call_tool_async( + tool, args=prepared_call.function_args, tool_context=prepared_call.tool_context, - ), + ) + + return await _execute_single_prepared_call( + invocation_context, prepared_call, agent, tool_runner=call_tool ) diff --git a/tests/unittests/flows/llm_flows/tools/test_functions_simple.py b/tests/unittests/flows/llm_flows/tools/test_functions_simple.py index bd8e6aaea91..fa38d29c30d 100644 --- a/tests/unittests/flows/llm_flows/tools/test_functions_simple.py +++ b/tests/unittests/flows/llm_flows/tools/test_functions_simple.py @@ -13,12 +13,16 @@ # limitations under the License. import asyncio +import contextvars +import threading from typing import Any from typing import Callable from unittest import mock from fastapi.openapi.models import HTTPBearer from google.adk.agents.llm_agent import Agent +from google.adk.agents.run_config import RunConfig +from google.adk.agents.run_config import ToolThreadPoolConfig from google.adk.auth.auth_tool import AuthConfig from google.adk.auth.auth_tool import AuthToolArguments from google.adk.events.event import Event @@ -941,6 +945,150 @@ async def yielding_async_function() -> dict: assert execution_order == ['sync_A', 'sync_B', 'async_C', 'async_D'] +@pytest.mark.asyncio +async def test_parallel_sync_function_calls_run_concurrently_with_tool_thread_pool(): + """Test that parallel calls to a sync function overlap with the pool set.""" + both_calls_running = threading.Barrier(2, timeout=5) + + def wait_for_other_call() -> dict: + both_calls_running.wait() + return {'result': 'done'} + + function_calls = [ + types.Part.from_function_call(name='wait_for_other_call', args={}), + types.Part.from_function_call(name='wait_for_other_call', args={}), + ] + function_responses = [ + types.Part.from_function_response( + name='wait_for_other_call', response={'result': 'done'} + ), + types.Part.from_function_response( + name='wait_for_other_call', response={'result': 'done'} + ), + ] + mock_model = testing_utils.MockModel.create( + responses=[function_calls, 'response1'] + ) + agent = Agent( + name='test_agent', model=mock_model, tools=[wait_for_other_call] + ) + runner = testing_utils.TestInMemoryRunner(agent) + + events = await runner.run_async_with_new_session( + 'test', RunConfig(tool_thread_pool_config=ToolThreadPoolConfig()) + ) + + assert testing_utils.simplify_events(events) == [ + ('test_agent', function_calls), + ('test_agent', function_responses), + ('test_agent', 'response1'), + ] + + +@pytest.mark.asyncio +async def test_sync_function_does_not_block_async_functions_with_tool_thread_pool(): + """Test that async functions run while a sync function blocks in the pool.""" + async_function_done = threading.Event() + + def blocking_sync_function() -> dict: + return {'async_function_done': async_function_done.wait(timeout=5)} + + async def yielding_async_function() -> dict: + await asyncio.sleep(0) + async_function_done.set() + return {'result': 'async_done'} + + function_calls = [ + types.Part.from_function_call(name='blocking_sync_function', args={}), + types.Part.from_function_call(name='yielding_async_function', args={}), + ] + mock_model = testing_utils.MockModel.create( + responses=[function_calls, 'response1'] + ) + agent = Agent( + name='test_agent', + model=mock_model, + tools=[blocking_sync_function, yielding_async_function], + ) + runner = testing_utils.TestInMemoryRunner(agent) + + events = await runner.run_async_with_new_session( + 'test', RunConfig(tool_thread_pool_config=ToolThreadPoolConfig()) + ) + + assert testing_utils.simplify_events(events)[1] == ( + 'test_agent', + [ + types.Part.from_function_response( + name='blocking_sync_function', + response={'async_function_done': True}, + ), + types.Part.from_function_response( + name='yielding_async_function', response={'result': 'async_done'} + ), + ], + ) + + +@pytest.mark.asyncio +async def test_sync_function_sees_callers_context_variables_with_tool_thread_pool(): + """Test that a sync function in the pool sees its caller's context variables.""" + request_id = contextvars.ContextVar('request_id', default=None) + + def read_request_id() -> dict: + return {'request_id': request_id.get()} + + mock_model = testing_utils.MockModel.create( + responses=[ + types.Part.from_function_call(name='read_request_id', args={}), + 'response1', + ] + ) + agent = Agent(name='test_agent', model=mock_model, tools=[read_request_id]) + runner = testing_utils.TestInMemoryRunner(agent) + request_id.set('request-1') + + events = await runner.run_async_with_new_session( + 'test', RunConfig(tool_thread_pool_config=ToolThreadPoolConfig()) + ) + + assert testing_utils.simplify_events(events)[1] == ( + 'test_agent', + types.Part.from_function_response( + name='read_request_id', response={'request_id': 'request-1'} + ), + ) + + +@pytest.mark.asyncio +async def test_async_function_runs_on_event_loop_thread_with_tool_thread_pool(): + """Test that an async function stays on the event loop thread with the pool set.""" + event_loop_thread = threading.get_ident() + + async def check_thread() -> dict: + return {'on_event_loop_thread': threading.get_ident() == event_loop_thread} + + mock_model = testing_utils.MockModel.create( + responses=[ + types.Part.from_function_call(name='check_thread', args={}), + 'response1', + ] + ) + agent = Agent(name='test_agent', model=mock_model, tools=[check_thread]) + runner = testing_utils.TestInMemoryRunner(agent) + + events = await runner.run_async_with_new_session( + 'test', RunConfig(tool_thread_pool_config=ToolThreadPoolConfig()) + ) + + assert testing_utils.simplify_events(events)[1] == ( + 'test_agent', + types.Part.from_function_response( + name='check_thread', response={'on_event_loop_thread': True} + ), + ) + + @pytest.mark.asyncio async def test_async_function_without_yield_blocks_others(): """Test that async functions without yield statements block other functions.""" diff --git a/tests/unittests/testing_utils.py b/tests/unittests/testing_utils.py index f308d0a2d1e..dd6899dc777 100644 --- a/tests/unittests/testing_utils.py +++ b/tests/unittests/testing_utils.py @@ -224,17 +224,23 @@ class TestInMemoryRunner(AfInMemoryRunner): """ async def run_async_with_new_session( - self, new_message: types.ContentUnion + self, + new_message: types.ContentUnion, + run_config: Optional[RunConfig] = None, ) -> list[Event]: collected_events: list[Event] = [] - async for event in self.run_async_with_new_session_agen(new_message): + async for event in self.run_async_with_new_session_agen( + new_message, run_config + ): collected_events.append(event) return collected_events async def run_async_with_new_session_agen( - self, new_message: types.ContentUnion + self, + new_message: types.ContentUnion, + run_config: Optional[RunConfig] = None, ) -> AsyncGenerator[Event, None]: session = await self.session_service.create_session( app_name='InMemoryRunner', user_id='test_user' @@ -243,6 +249,7 @@ async def run_async_with_new_session_agen( user_id=session.user_id, session_id=session.id, new_message=get_user_content(new_message), + run_config=run_config, ) async with Aclosing(agen): async for event in agen: From 643df966a0fa223454e34a4eee1eddd68931ffa3 Mon Sep 17 00:00:00 2001 From: Aarav Mittal <137450929+a2105z@users.noreply.github.com> Date: Tue, 29 Sep 2026 10:29:05 -0700 Subject: [PATCH 09/29] fix: drop unpairable trailing FRs in rearrange Merge https://github.com/google/adk-python/pull/6752 Fixes #6751 PiperOrigin-RevId: 990390762 --- .../adk/flows/llm_flows/tools/_rearranger.py | 229 ++++++---- .../flows/llm_flows/tools/test_rearranger.py | 428 +++++++++++++++++- 2 files changed, 554 insertions(+), 103 deletions(-) diff --git a/src/google/adk/flows/llm_flows/tools/_rearranger.py b/src/google/adk/flows/llm_flows/tools/_rearranger.py index 5e3e1ff2094..5ed06d928b5 100644 --- a/src/google/adk/flows/llm_flows/tools/_rearranger.py +++ b/src/google/adk/flows/llm_flows/tools/_rearranger.py @@ -22,7 +22,6 @@ from google.genai import types from ....events.event import Event -from ._functions import _collect_function_call_ids logger = logging.getLogger('google_adk.' + __name__) @@ -181,8 +180,7 @@ def drop_orphaned_function_responses( outcome the same wherever it appears, and keeps unpaired results from being forwarded to providers that reject them. - Responses without an id are left alone: ids are stripped on the way out for - some model families, so a missing id does not imply a missing call. + Responses without a preceding matching function_call are dropped as orphans. Args: events: The events being assembled into request contents. @@ -190,31 +188,44 @@ def drop_orphaned_function_responses( Returns: The events with orphaned function_response parts removed. """ - call_ids = _collect_function_call_ids(events) - + seen_call_ids: set[str] = set() + seen_idless_call_names: set[str] = set() orphaned_ids: list[str] = [] result_events: list[Event] = [] for event in events: parts = event.content.parts if event.content else None - if not parts or not event.get_function_responses(): + if parts and event.get_function_responses(): + kept_parts: list[types.Part] = [] + for part in parts: + response = part.function_response + if response: + is_matched = ( + response.id in seen_call_ids + if response.id + else ( + bool(response.name) + and response.name in seen_idless_call_names + ) + ) + if not is_matched: + orphaned_ids.append(response.id or '') + continue + kept_parts.append(part) + + if kept_parts: + if len(kept_parts) != len(parts): + event = event.model_copy(deep=True) + if event.content: + event.content.parts = kept_parts + result_events.append(event) + else: result_events.append(event) - continue - kept_parts: list[types.Part] = [] - for part in parts: - response = part.function_response - if response and response.id and response.id not in call_ids: - orphaned_ids.append(response.id) - continue - kept_parts.append(part) - - if not kept_parts: - continue - if len(kept_parts) != len(parts): - event = event.model_copy(deep=True) - if event.content: - event.content.parts = kept_parts - result_events.append(event) + for fc in event.get_function_calls(): + if fc.id: + seen_call_ids.add(fc.id) + elif fc.name: + seen_idless_call_names.add(fc.name) if orphaned_ids: logger.warning( @@ -308,6 +319,25 @@ def drop_orphaned_function_calls( return result_events +def _find_owning_call_event_index( + history_events: list[Event], + response: types.FunctionResponse, +) -> int: + for idx in range(len(history_events) - 1, -1, -1): + if any( + (bool(response.id) and c.id == response.id) + or ( + not response.id + and not c.id + and bool(response.name) + and c.name == response.name + ) + for c in history_events[idx].get_function_calls() + ): + return idx + return -1 + + def rearrange_events_for_latest_function_response( events: list[Event], ) -> list[Event]: @@ -317,87 +347,104 @@ def rearrange_events_for_latest_function_response( between the initial function_call and the latest function_response will be removed. + If the latest event carries function responses with no matching function + call in history (an orphaned FR), those responses are dropped and history + is rearranged from the remaining events. + Args: events: A list of events. Returns: A list of events with the latest function_response rearranged. """ - if len(events) < 2: - # No need to process, since there is no function_call. + events = drop_orphaned_function_responses(events) + if len(events) < 2 or not events[-1].get_function_responses(): return events - function_responses = events[-1].get_function_responses() - if not function_responses: - # No need to process, since the latest event is not function_response. + trailing = events[-1] + parts = trailing.content.parts if trailing.content else None + if not parts: return events - function_responses_ids = set() - for function_response in function_responses: - function_responses_ids.add(function_response.id) - - function_calls = events[-2].get_function_calls() - - if function_calls: - for function_call in function_calls: - # The latest function_response is already matched - if function_call.id in function_responses_ids: - return events - - function_call_event_idx = -1 - # look for corresponding function call event reversely - for idx in range(len(events) - 2, -1, -1): - event = events[idx] - function_calls = event.get_function_calls() - if function_calls: - for function_call in function_calls: - if function_call.id in function_responses_ids: - function_call_event_idx = idx - function_call_ids = { - function_call.id for function_call in function_calls - } - # last response event should only contain the responses for the - # function calls in the same function call event - if not function_responses_ids.issubset(function_call_ids): - raise ValueError( - 'Last response event should only contain the responses for the' - ' function calls in the same function call event. Function' - f' call ids found : {function_call_ids}, function response' - f' ids provided: {function_responses_ids}' - ) - # collect all function responses from the function call event to - # the last response event - function_responses_ids = function_call_ids - break - - if function_call_event_idx == -1: - logger.debug( - 'No function call event found for function responses ids: %s in' - ' event list: %s', - function_responses_ids, - events, - ) - raise ValueError( - 'No function call event found for function responses ids:' - f' {function_responses_ids}' + history_events = events[:-1] + parts_by_call_idx: dict[int, list[types.Part]] = {} + non_fr_parts: list[types.Part] = [] + for part in parts: + if part.function_response: + call_idx = _find_owning_call_event_index( + history_events, part.function_response + ) + if call_idx != -1: + parts_by_call_idx.setdefault(call_idx, []).append(part) + else: + non_fr_parts.append(part) + + if not parts_by_call_idx: + return events + + latest_call_idx = max(parts_by_call_idx.keys()) + parts_by_call_idx[latest_call_idx].extend(non_fr_parts) + + if latest_call_idx == len(events) - 2 and len(parts_by_call_idx) == 1: + return events + + def _merged_response_for_call(call_idx: int) -> Event: + split_event = trailing.model_copy(deep=True) + if split_event.content: + split_event.content.parts = parts_by_call_idx[call_idx] + intermediate: list[Event] = [] + for ev_idx in range(call_idx + 1, len(events) - 1): + ev = events[ev_idx] + ev_parts = ev.content.parts if ev.content else None + if not ev_parts or not ev.get_function_responses(): + continue + is_call_event = bool(ev.get_function_calls()) + matched_parts = [ + p + for p in ev_parts + if ( + _find_owning_call_event_index( + events[:ev_idx], p.function_response + ) + == call_idx + if p.function_response + else not is_call_event + ) + ] + if any(p.function_response for p in matched_parts): + if len(matched_parts) != len(ev_parts): + ev = ev.model_copy(deep=True) + if ev.content: + ev.content.parts = matched_parts + intermediate.append(ev) + all_resps = intermediate + [split_event] + return ( + merge_function_response_events(all_resps) + if len(all_resps) > 1 + else all_resps[0] ) - # collect all function response between last function response event - # and function call event - - function_response_events: list[Event] = [] - for idx in range(function_call_event_idx + 1, len(events) - 1): - event = events[idx] - function_responses = event.get_function_responses() - if function_responses and any([ - function_response.id in function_responses_ids - for function_response in function_responses - ]): - function_response_events.append(event) - function_response_events.append(events[-1]) - - result_events = events[: function_call_event_idx + 1] - result_events.append(merge_function_response_events(function_response_events)) + result_events: list[Event] = [] + for idx in range(latest_call_idx + 1): + ev = events[idx] + ev_parts = ev.content.parts if ev.content else None + if ev_parts and ev.get_function_responses(): + kept_parts = [ + p + for p in ev_parts + if not p.function_response + or _find_owning_call_event_index(events[:idx], p.function_response) + not in parts_by_call_idx + ] + if not any(p.function_response or p.function_call for p in kept_parts): + continue + if len(kept_parts) != len(ev_parts): + ev = ev.model_copy(deep=True) + if ev.content: + ev.content.parts = kept_parts + result_events.append(ev) + if idx in parts_by_call_idx: + result_events.append(_merged_response_for_call(idx)) return result_events diff --git a/tests/unittests/flows/llm_flows/tools/test_rearranger.py b/tests/unittests/flows/llm_flows/tools/test_rearranger.py index d9399b34276..4ae1905ecb5 100644 --- a/tests/unittests/flows/llm_flows/tools/test_rearranger.py +++ b/tests/unittests/flows/llm_flows/tools/test_rearranger.py @@ -25,6 +25,7 @@ from google.adk.flows.llm_flows.tools._rearranger import merge_function_response_events from google.adk.flows.llm_flows.tools._rearranger import rearrange_events_for_async_function_responses_in_history from google.adk.flows.llm_flows.tools._rearranger import rearrange_events_for_latest_function_response +from google.adk.models.anthropic_llm import content_to_message_param from google.genai import types import pytest @@ -65,7 +66,7 @@ def _resp_event( def test_drop_orphaned_responses_prunes_unpaired_and_preserves_valid(): - """Unpaired function response IDs are pruned while matched and ID-less responses survive.""" + """Unpaired and ID-less function responses are pruned while matched responses survive.""" call = _call_event("c1", "lookup") valid_resp = _resp_event("c1", "lookup", "found") no_id_resp = _resp_event(None, "legacy", "ok") @@ -74,7 +75,7 @@ def test_drop_orphaned_responses_prunes_unpaired_and_preserves_valid(): result = drop_orphaned_function_responses(events) - assert result == [call, valid_resp, no_id_resp] + assert result == [call, valid_resp] def test_drop_orphaned_responses_removes_event_when_all_parts_orphaned(): @@ -221,8 +222,6 @@ def test_drop_orphaned_calls_prunes_unanswered_in_parallel_tool_calls(): def test_drop_orphaned_calls_prevents_unclosed_tool_use_in_anthropic_conversion(): """Pruned events converted for Anthropic contain no unclosed tool_use blocks.""" - from google.adk.models.anthropic_llm import content_to_message_param - orphan_call = _call_event("fc_interrupted", "slow_tool") user_turn = Event( author="user", @@ -371,15 +370,189 @@ def test_rearrange_latest_response_moves_to_call_and_prunes_intervening(): } -def test_rearrange_latest_response_missing_matching_call_raises_value_error(): - """A trailing response with no matching preceding call raises ValueError.""" - events = [ - Event(author="user", content=types.UserContent("hello")), - _resp_event("missing_call_id"), - ] +@pytest.mark.parametrize( + "resp_event", + [ + _resp_event("missing_id"), + _resp_event(None, "missing_tool"), + _resp_event("", "missing_tool"), + ], + ids=["unmatched-id", "idless", "empty-id"], +) +def test_rearrange_latest_response_drops_orphans( + *, + resp_event: Event, +) -> None: + """Trailing FR with unpairable, None, or empty id is dropped.""" + user_msg = Event(author="user", content=types.UserContent("hello")) + result = rearrange_events_for_latest_function_response([user_msg, resp_event]) + assert result == [user_msg] + + +def test_rearrange_latest_response_drops_orphan_part_preserves_valid() -> None: + """An unmatched FR part is dropped while a matched sibling part is kept.""" + call = _call_event("c1", "tool_a") + trailing = Event( + author="user", + content=types.UserContent([ + types.Part( + function_response=types.FunctionResponse( + id="c1", name="tool_a", response={"ok": True} + ) + ), + types.Part( + function_response=types.FunctionResponse( + id="extra", name="tool_b", response={"ok": True} + ) + ), + ]), + ) + result = rearrange_events_for_latest_function_response([call, trailing]) + assert len(result) == 2 + assert [r.id for r in result[1].get_function_responses()] == ["c1"] + + +def test_rearrange_latest_response_preserves_non_fr_parts_when_orphan_dropped() -> ( + None +): + """Non-FR parts in the trailing event are preserved when an orphan is dropped.""" + user_msg = Event(author="user", content=types.UserContent("hello")) + trailing = Event( + author="user", + content=types.UserContent([ + types.Part(text="keep this text"), + types.Part( + function_response=types.FunctionResponse( + name="ghost", response={"err": 1} + ) + ), + ]), + ) + result = rearrange_events_for_latest_function_response([user_msg, trailing]) + assert len(result) == 2 + assert result[0] == user_msg + assert result[1].content is not None and result[1].content.parts is not None + assert result[1].content.parts[0].text == "keep this text" + + +def test_rearrange_latest_response_splits_across_calls() -> None: + """Trailing responses for multiple calls split and pair per owning call.""" + call_slow = _call_event("c_slow", "slow_tool") + user_wait = Event(author="user", content=types.UserContent("waiting...")) + call_fast = _call_event("c_fast", "fast_tool") + trailing = Event( + author="user", + content=types.UserContent([ + types.Part( + function_response=types.FunctionResponse( + id="c_fast", name="fast_tool", response={"fast": True} + ) + ), + types.Part( + function_response=types.FunctionResponse( + id="c_slow", name="slow_tool", response={"slow": True} + ) + ), + ]), + ) + events = [call_slow, user_wait, call_fast, trailing] + result = rearrange_events_for_latest_function_response(events) + + assert len(result) == 5 + assert result[0] == call_slow + assert [r.name for r in result[1].get_function_responses()] == ["slow_tool"] + assert result[2] == user_wait + assert result[3] == call_fast + assert [r.name for r in result[4].get_function_responses()] == ["fast_tool"] + + +def test_rearrange_latest_response_merges_intermediate_for_earlier_call() -> ( + None +): + """Trailing response for an earlier call merges with and updates its intermediate response.""" + call_slow = _call_event("c_slow", "slow_tool") + progress_slow = _resp_event("c_slow", "slow_tool", {"progress": "50%"}) + call_fast = _call_event("c_fast", "fast_tool") + trailing = Event( + author="user", + content=types.UserContent([ + types.Part( + function_response=types.FunctionResponse( + id="c_fast", name="fast_tool", response={"fast": True} + ) + ), + types.Part( + function_response=types.FunctionResponse( + id="c_slow", + name="slow_tool", + response={"progress": "100%", "result": "final"}, + ) + ), + ]), + ) + events = [call_slow, progress_slow, call_fast, trailing] + result = rearrange_events_for_latest_function_response(events) + + assert len(result) == 4 + assert result[0] == call_slow + assert result[1].get_function_responses()[0].response == { + "progress": "100%", + "result": "final", + } + assert result[2] == call_fast + assert [r.name for r in result[3].get_function_responses()] == ["fast_tool"] + + +def test_rearrange_latest_response_splits_shared_intermediate_progress_per_call() -> ( + None +): + """A progress event covering two calls only merges each call's own part into its split.""" + call_slow = _call_event("c_slow", "slow_tool") + call_fast = _call_event("c_fast", "fast_tool") + progress_both = Event( + author="user", + content=types.UserContent([ + types.Part( + function_response=types.FunctionResponse( + id="c_slow", name="slow_tool", response={"progress": "50%"} + ) + ), + types.Part( + function_response=types.FunctionResponse( + id="c_fast", name="fast_tool", response={"progress": "50%"} + ) + ), + ]), + ) + trailing = Event( + author="user", + content=types.UserContent([ + types.Part( + function_response=types.FunctionResponse( + id="c_fast", + name="fast_tool", + response={"progress": "100%", "fast": True}, + ) + ), + types.Part( + function_response=types.FunctionResponse( + id="c_slow", + name="slow_tool", + response={"progress": "100%", "result": "final"}, + ) + ), + ]), + ) + events = [call_slow, call_fast, progress_both, trailing] + result = rearrange_events_for_latest_function_response(events) - with pytest.raises(ValueError, match="No function call event found"): - rearrange_events_for_latest_function_response(events) + assert len(result) == 4 + assert [(r.id, r.response) for r in result[1].get_function_responses()] == [ + ("c_slow", {"progress": "100%", "result": "final"}) + ] + assert [(r.id, r.response) for r in result[3].get_function_responses()] == [ + ("c_fast", {"progress": "100%", "fast": True}) + ] def test_rearrange_history_reused_id_across_tools_pairs_correctly(): @@ -490,3 +663,234 @@ def test_backward_compatibility_aliases_exported(): _tool_call_rearranger._rearrange_events_for_latest_function_response is rearrange_events_for_latest_function_response ) + + +def test_rearrange_latest_response_merges_intermediate_response_after_intervening_call() -> ( + None +): + """Intermediate response for an earlier call merges when positioned after a later call.""" + call_parallel = Event( + author="test_agent", + content=types.Content( + role="model", + parts=[ + types.Part( + function_call=types.FunctionCall( + id="c1", name="tool_1", args={} + ) + ), + types.Part( + function_call=types.FunctionCall( + id="c2", name="tool_2", args={} + ) + ), + ], + ), + ) + call_other = _call_event("c3", "tool_3") + resp_c1 = _resp_event("c1", "tool_1", {"r1": True}) + trailing = Event( + author="user", + content=types.UserContent([ + types.Part( + function_response=types.FunctionResponse( + id="c2", name="tool_2", response={"r2": True} + ) + ), + types.Part( + function_response=types.FunctionResponse( + id="c3", name="tool_3", response={"r3": True} + ) + ), + ]), + ) + events = [call_parallel, call_other, resp_c1, trailing] + result = rearrange_events_for_latest_function_response(events) + + assert len(result) == 4 + assert result[0] == call_parallel + resp_ids = [r.id for r in result[1].get_function_responses()] + assert "c1" in resp_ids + assert "c2" in resp_ids + + +def test_rearrange_latest_response_reused_id_preserves_earlier_settled_response() -> ( + None +): + """A reused call ID with intervening turn preserves the earlier call's settled response.""" + call1 = _call_event("call_1", "tool_a") + resp1 = _resp_event("call_1", "tool_a", "first") + call2 = _call_event("call_1", "tool_b") + intervening = Event(author="user", content=types.UserContent("intervening")) + resp2 = _resp_event("call_1", "tool_b", "second") + events = [call1, resp1, call2, intervening, resp2] + + result = rearrange_events_for_latest_function_response(events) + + assert len(result) == 4 + assert result[0] == call1 + assert result[1] == resp1 + assert result[2] == call2 + assert result[3].get_function_responses()[0].response == {"result": "second"} + + +def test_drop_orphaned_responses_preserves_idless_response_for_idless_call() -> ( + None +): + """ID-less function responses are preserved when preceded by a matching call.""" + call = Event( + author="test_agent", + content=types.Content( + role="model", + parts=[types.Part(function_call=types.FunctionCall(name="legacy"))], + ), + ) + resp = Event( + author="user", + content=types.UserContent([ + types.Part( + function_response=types.FunctionResponse( + name="legacy", response={"found": True} + ) + ) + ]), + ) + events = [call, resp] + + result = drop_orphaned_function_responses(events) + + assert result == [call, resp] + + +def test_find_owning_call_event_index_does_not_match_different_idless_tools() -> ( + None +): + """An ID-less response must not match an ID-less call for a different tool.""" + call = Event( + author="test_agent", + content=types.Content( + role="model", + parts=[types.Part(function_call=types.FunctionCall(name="weather"))], + ), + ) + resp = types.FunctionResponse(name="calculator", response={"and": 42}) + assert _tool_call_rearranger._find_owning_call_event_index([call], resp) == -1 + + +def test_rearrange_latest_response_merges_intermediate_for_earlier_idless_call() -> ( + None +): + """Trailing response for an earlier ID-less call merges its intermediate response.""" + call_slow = Event( + author="test_agent", + content=types.Content( + role="model", + parts=[ + types.Part(function_call=types.FunctionCall(name="slow_tool")) + ], + ), + ) + progress_slow = Event( + author="user", + content=types.UserContent([ + types.Part( + function_response=types.FunctionResponse( + name="slow_tool", response={"progress": "50%"} + ) + ) + ]), + ) + call_fast = Event( + author="test_agent", + content=types.Content( + role="model", + parts=[ + types.Part(function_call=types.FunctionCall(name="fast_tool")) + ], + ), + ) + trailing = Event( + author="user", + content=types.UserContent([ + types.Part( + function_response=types.FunctionResponse( + name="fast_tool", response={"fast": True} + ) + ), + types.Part( + function_response=types.FunctionResponse( + name="slow_tool", + response={"progress": "100%", "result": "final"}, + ) + ), + ]), + ) + events = [call_slow, progress_slow, call_fast, trailing] + result = rearrange_events_for_latest_function_response(events) + + assert len(result) == 4 + assert result[0] == call_slow + assert result[1].get_function_responses()[0].response == { + "progress": "100%", + "result": "final", + } + assert result[2] == call_fast + assert [r.name for r in result[3].get_function_responses()] == ["fast_tool"] + + +def test_rearrange_latest_response_preserves_call_event_carrying_consumed_response() -> ( + None +): + """Call event with a consumed response keeps its call and response.""" + call_slow = _call_event("c_slow", "slow_tool") + call_fast_with_progress = Event( + author="test_agent", + content=types.Content( + role="model", + parts=[ + types.Part( + function_response=types.FunctionResponse( + id="c_slow", + name="slow_tool", + response={"progress": "50%"}, + ) + ), + types.Part( + function_call=types.FunctionCall( + id="c_fast", name="fast_tool", args={} + ) + ), + ], + ), + ) + trailing = Event( + author="user", + content=types.UserContent([ + types.Part( + function_response=types.FunctionResponse( + id="c_fast", name="fast_tool", response={"fast": True} + ) + ), + types.Part( + function_response=types.FunctionResponse( + id="c_slow", + name="slow_tool", + response={"progress": "100%", "result": "final"}, + ) + ), + ]), + ) + events = [call_slow, call_fast_with_progress, trailing] + result = rearrange_events_for_latest_function_response(events) + + assert len(result) == 4 + assert result[0] == call_slow + assert result[1].get_function_calls() == [] + assert [(r.id, r.response) for r in result[1].get_function_responses()] == [ + ("c_slow", {"progress": "100%", "result": "final"}) + ] + assert [c.id for c in result[2].get_function_calls()] == ["c_fast"] + assert result[2].get_function_responses() == [] + assert [(r.id, r.response) for r in result[3].get_function_responses()] == [ + ("c_fast", {"fast": True}) + ] From 7298e09acb485c4b226ec552c899f661df64fe49 Mon Sep 17 00:00:00 2001 From: George Weale Date: Tue, 29 Sep 2026 10:55:21 -0700 Subject: [PATCH 10/29] fix: replay a parallel tool call that never ran when a sibling answered When a resumed model turn had parallel calls and only some had a response, the resume decision treated one answer as covering the whole turn, so it continued to the model and the call that never ran was dropped. The decision now replays just the calls with no response, and only when the agent has written nothing since the call, so answered calls do not run a second time. Close #7108 Co-authored-by: George Weale PiperOrigin-RevId: 990408003 --- .../adk/flows/llm_flows/core/_resume.py | 50 ++++++++++- .../flows/llm_flows/core/test_resume.py | 85 +++++++++++++++++++ .../runners/test_resume_invocation.py | 68 +++++++++++++++ 3 files changed, 201 insertions(+), 2 deletions(-) diff --git a/src/google/adk/flows/llm_flows/core/_resume.py b/src/google/adk/flows/llm_flows/core/_resume.py index ffcb9cb78ce..1106808430c 100644 --- a/src/google/adk/flows/llm_flows/core/_resume.py +++ b/src/google/adk/flows/llm_flows/core/_resume.py @@ -49,7 +49,11 @@ class ResumeAction(enum.Enum): """A tool call is still unanswered; stop without emitting anything.""" REPLAY_CALLS = 'replay_calls' - """A tool call was never executed; run the calls on `ResumeDecision.event`.""" + """A tool call was never executed; run the calls on `ResumeDecision.event`. + + That event may be a copy holding only the unexecuted calls, so it is not + always one of `session.events` and must not be compared by identity. + """ @dataclasses.dataclass(frozen=True) @@ -207,6 +211,42 @@ def _needs_call_replay( ) +def _unexecuted_calls_event( + call_event: Event, later_events: list[Event] +) -> Event | None: + """`call_event` cut down to the calls that never ran, or None if all did. + + A call ran when a response in `later_events` carries its id, or its name + with no id. None is also returned once the agent has written any event with + content after `call_event`: tool responses and auth or confirmation requests + are only written after the whole batch ran, so a call still missing its + response then ran and lost it, or is pending. + """ + if any( + ev.author == call_event.author and ev.content is not None + for ev in later_events + ): + return None + responses = [fr for ev in later_events for fr in ev.get_function_responses()] + answered_ids = {fr.id for fr in responses if fr.id is not None} + answered_names = {fr.name for fr in responses if fr.id is None} + unexecuted_ids = { + fc.id + for fc in call_event.get_function_calls() + if fc.id not in answered_ids and fc.name not in answered_names + } + if not unexecuted_ids or call_event.content is None: + return None + parts = [ + part + for part in call_event.content.parts or [] + if part.function_call is None or part.function_call.id in unexecuted_ids + ] + return call_event.model_copy( + update={'content': call_event.content.model_copy(update={'parts': parts})} + ) + + def decide_resume( invocation_context: InvocationContext, events: list[Event], @@ -222,7 +262,9 @@ def decide_resume( Returns: PAUSE when a call is still unanswered, REPLAY_CALLS (naming the event whose - calls to run) when a call was never executed, else CONTINUE. + calls to run, which may be a copy holding only the unexecuted calls rather + than one of `session.events`) when a call was never executed, else + CONTINUE. """ paused_by_last = invocation_context.should_pause_invocation(events[-1]) if not paused_by_last and _pause_left_calls_unanswered( @@ -270,6 +312,10 @@ def decide_resume( pause = True elif _needs_call_replay(call_names, answers, from_sub_branch): return ResumeDecision(ResumeAction.REPLAY_CALLS, call_event) + elif unexecuted := _unexecuted_calls_event( + call_event, events[call_idx + 1 :] + ): + return ResumeDecision(ResumeAction.REPLAY_CALLS, unexecuted) return ResumeDecision(ResumeAction.PAUSE if pause else ResumeAction.CONTINUE) diff --git a/tests/unittests/flows/llm_flows/core/test_resume.py b/tests/unittests/flows/llm_flows/core/test_resume.py index afebde31320..6ac2c4e8104 100644 --- a/tests/unittests/flows/llm_flows/core/test_resume.py +++ b/tests/unittests/flows/llm_flows/core/test_resume.py @@ -29,6 +29,7 @@ from google.adk.flows.llm_flows.core._resume import decide_step_resume from google.adk.flows.llm_flows.core._resume import ResumeAction from google.adk.flows.llm_flows.core._resume import ResumeDecision +from google.adk.flows.llm_flows.functions import REQUEST_EUC_FUNCTION_CALL_NAME from google.adk.workflow.utils._workflow_hitl_utils import REQUEST_INPUT_FUNCTION_CALL_NAME from google.genai import types import pytest @@ -76,6 +77,28 @@ def _response_event( ) +def _parallel_call_event(calls: list[tuple[str, str]]) -> Event: + return Event( + author='agent', + invocation_id='inv-1', + content=types.Content( + role='model', + parts=[ + types.Part( + function_call=types.FunctionCall( + id=call_id, name=name, args={} + ) + ) + for name, call_id in calls + ], + ), + ) + + +def _replayed_ids(decision: ResumeDecision) -> list[str | None]: + return [fc.id for fc in decision.replay_event().get_function_calls()] + + def _text_event(text: str) -> Event: return Event( author='agent', @@ -321,6 +344,55 @@ def test_parallel_calls_all_answered_continue(self): ) assert decision.action is ResumeAction.CONTINUE + @pytest.mark.parametrize( + 'sibling_name', ['fetch', 'ask'], ids=['other_name', 'same_name'] + ) + def test_parallel_call_that_never_ran_is_replayed_alone(self, sibling_name): + call = _parallel_call_event([('ask', 'c1'), (sibling_name, 'c2')]) + events = [call, _response_event('ask', 'c1')] + decision = decide_resume( + self._ctx(), events, {'ask': object(), 'fetch': object()} + ) + assert decision.action is ResumeAction.REPLAY_CALLS + assert _replayed_ids(decision) == ['c2'] + + def test_response_without_an_id_answers_its_call_by_name(self): + call = _parallel_call_event([('ask', 'c1'), ('fetch', 'c2')]) + events = [call, _response_event('ask', None)] + decision = decide_resume( + self._ctx(), events, {'ask': object(), 'fetch': object()} + ) + assert _replayed_ids(decision) == ['c2'] + + def test_sibling_missing_a_response_after_an_auth_resume_is_not_replayed( + self, + ): + auth_request = Event( + author='agent', + invocation_id='inv-1', + long_running_tool_ids={'a1'}, + content=types.Content( + role='user', + parts=[ + types.Part( + function_call=types.FunctionCall( + id='a1', name=REQUEST_EUC_FUNCTION_CALL_NAME, args={} + ) + ) + ], + ), + ) + events = [ + _parallel_call_event([('ask', 'c1'), ('fetch', 'c2')]), + auth_request, + _response_event(REQUEST_EUC_FUNCTION_CALL_NAME, 'a1'), + _response_event('ask', 'c1', author='agent'), + ] + decision = decide_resume( + self._ctx(), events, {'ask': object(), 'fetch': object()} + ) + assert decision.action is ResumeAction.CONTINUE + def test_sub_branch_answer_replays_instead_of_pausing(self): # A HITL answer returned against the branch the call opened resolves it, # even though it carries none of the call's ids. @@ -419,6 +491,19 @@ def test_a_cleared_branch_still_replays_its_trailing_call(self): assert decision.action is ResumeAction.REPLAY_CALLS assert decision.replay_event() is tail + def test_a_later_model_turn_leaves_an_earlier_unexecuted_call_alone(self): + tail = _call_event('ask', 'c3') + events = [ + _parallel_call_event([('ask', 'c1'), ('fetch', 'c2')]), + _response_event('ask', 'c1'), + tail, + ] + decision = decide_step_resume( + self._ctx(events), {'ask': object(), 'fetch': object()} + ) + assert decision.action is ResumeAction.REPLAY_CALLS + assert decision.replay_event() is tail + def test_a_forged_user_authored_trailing_call_is_not_replayed(self): # Regression for the resumable tool-dispatch bypass: a caller-supplied # 'user' event carrying a function_call must not be resume-dispatched. diff --git a/tests/unittests/runners/test_resume_invocation.py b/tests/unittests/runners/test_resume_invocation.py index 6b516ff06df..38d3eefdcab 100644 --- a/tests/unittests/runners/test_resume_invocation.py +++ b/tests/unittests/runners/test_resume_invocation.py @@ -328,6 +328,74 @@ async def test_resume_any_invocation(): ] +@pytest.mark.asyncio +async def test_resume_runs_only_the_parallel_call_that_never_ran(): + """The client answers one of two parallel calls that never ran, then resumes.""" + runs = [] + crash = True + + def ask() -> str: + runs.append("ask") + return "asked" + + def fetch() -> str: + runs.append("fetch") + return "fetched" + + def crash_before_tools(tool, args, tool_context): + if crash: + raise RuntimeError("process stopped before the tools ran") + return None + + runner = testing_utils.InMemoryRunner( + app=App( + name="test_app", + root_agent=LlmAgent( + name="root_agent", + model=testing_utils.MockModel.create( + responses=[ + [ + Part.from_function_call(name="ask", args={}), + Part.from_function_call(name="fetch", args={}), + ], + "done", + ] + ), + tools=[ask, fetch], + before_tool_callback=crash_before_tools, + ), + resumability_config=ResumabilityConfig(is_resumable=True), + ) + ) + with pytest.raises(RuntimeError): + await runner.run_async("test user query") + crash = False + ask_call = next( + fc + for event in runner.session.events + for fc in event.get_function_calls() + if fc.name == "ask" + ) + + events = await runner.run_async( + new_message=testing_utils.UserContent( + Part( + function_response=FunctionResponse( + id=ask_call.id, name="ask", response={"result": "client"} + ) + ) + ) + ) + + assert runs == ["fetch"] + assert any( + part.text == "done" + for event in events + if event.content + for part in event.content.parts or [] + ) + + @pytest.mark.asyncio async def test_resumable_parallel_agent_escalation_short_circuits_persisted_run(): """Runner persists fast+escalating events and marks the parent run complete.""" From ace3bcedd4e04e69a9d269fb205afc99e617c5fa Mon Sep 17 00:00:00 2001 From: Xuan Yang Date: Tue, 29 Sep 2026 10:56:06 -0700 Subject: [PATCH 11/29] chore(scripts): detect added files through git only The new-file check read added files and commit messages from several version control systems. Contributors and CI use git through pre-commit, and jj's default colocated repositories already take the git path, so the rest only added code to maintain. Outside a git work tree the check now reports that it could not determine the added files. A `//`-prefixed .py path with no file behind it, the form Perforce depot paths take, is now refused. Co-authored-by: Xuan Yang PiperOrigin-RevId: 990408500 --- scripts/check_new_py_files.py | 430 ++++++---------- .../scripts/test_check_new_py_files.py | 484 +++++------------- 2 files changed, 283 insertions(+), 631 deletions(-) diff --git a/scripts/check_new_py_files.py b/scripts/check_new_py_files.py index 68107cd42df..5e1db32cd68 100644 --- a/scripts/check_new_py_files.py +++ b/scripts/check_new_py_files.py @@ -33,7 +33,7 @@ Modes for finding added files: - Baseline Diff Mode (CI): python scripts/check_new_py_files.py --baseline-dir /path/to/origin-main -- VCS Detection Mode (Local / Pre-commit): +- Git Detection Mode (Local / Pre-commit): python scripts/check_new_py_files.py - Explicit File List: python scripts/check_new_py_files.py file1.py file2.py @@ -63,28 +63,12 @@ _PACKAGE_RELPATH = os.path.join('src', 'google', 'adk') _DOCS_GUIDES_RELPATH = os.path.join('docs', 'guides') -# The package's import path, i.e. _PACKAGE_RELPATH without the src/ root that -# only this checkout uses. A repository laying the package out differently -# still ends its path to it with these components. -_PACKAGE_IMPORT_PATH = 'google/adk' - _EXIT_OK = 0 _EXIT_VIOLATIONS = 1 _EXIT_SETUP_ERROR = 2 _EXIT_INDETERMINATE = 3 -# The newest revision this checkout shares with the server, as a Mercurial -# revset: everything after it is the change under construction. Revisions the -# server already has are in the public phase and local work is draft, so this -# is the base the change sits on; with no local commits it evaluates to '.'. -_SYNCED_BASE = 'last(public() & ::.)' - -# The local commits themselves, as a Mercurial revset: the ones after -# `_SYNCED_BASE`, i.e. the work the change is made of. Used to read the commit -# messages over the same range the added-file scan covers. -_LOCAL_COMMITS = 'draft() & ::.' - -# The commit range a git checkout's HEAD covers when nothing is staged. On a +# The commit range a git work tree's HEAD covers when nothing is staged. On a # pull request this is the base branch to the merge commit, i.e. the pull # request's own commits, which is both the file set to check and the place a # NO_UNIT_GUIDE waiver would be written. @@ -135,8 +119,9 @@ 'If a unit guide is not required for this file, explain why with a' " 'NO_UNIT_GUIDE=' tag in the commit message of the change that" ' adds it. Where no message can be read, as when the change is only' - ' staged or the tree carries no version control, set the tag in the' - " environment instead: NO_UNIT_GUIDE='' git commit ...\n" + ' staged or the tree is not a git work tree, set the tag in the' + ' environment instead; for a staged change, that is' + " NO_UNIT_GUIDE='' git commit.\n" 'See .agents/skills/adk-unit-guide/SKILL.md for details on creating unit' ' guides.' ) @@ -367,214 +352,106 @@ def _git_is_mid_commit(root: str) -> bool | None: return bool(staged_any.strip()) -def get_vcs_added_files(root: str = '.') -> set[str] | None: - """Detects added files using local VCS (git, jj, hg, g4, p4). +def _in_git_work_tree(root: str) -> bool: + """Whether git is installed and root lies inside a git work tree.""" + if not shutil.which('git'): + return False + # The exit status alone is not enough: inside the .git directory itself the + # command succeeds and prints `false`. + code, out = _run_cmd(['git', 'rev-parse', '--is-inside-work-tree'], cwd=root) + return code == 0 and out == 'true' + + +def get_git_added_files(root: str = '.') -> set[str] | None: + """Detects the files a change adds, using git. + + The change is what is staged, when anything is, and otherwise the range + HEAD~1..HEAD, which on a pull request covers the contributor's commits. A + rename counts only when it exposes a name no rule has judged. Args: root: The root directory of the repository. Returns: - A set of added file paths if a supported VCS is detected, or None if no - supported VCS was detected. + The added file paths, or None when git is not installed, root is not in a + git work tree, or git could not say what the change adds. """ - # 1. git - if shutil.which('git'): - code, _ = _run_cmd(['git', 'rev-parse', '--is-inside-work-tree'], cwd=root) - if code == 0: - mid_commit = _git_is_mid_commit(root) - if mid_commit is None: - print( - 'git is active but its index cannot be read, so whether this' - ' change is staged or committed is unknown, and so is what it' - ' adds.', - file=sys.stderr, - ) - return None - if mid_commit: - return _git_added_paths(['git', 'diff', '--cached'], root) - range_code, head_diff = _run_cmd( - ['git', 'diff', _GIT_HEAD_RANGE, '--name-only', _GIT_ADD_FILTER], - cwd=root, - ) - if range_code != 0: - # HEAD~1 is unreachable, as in a depth-1 clone. The range resolved to - # nothing rather than to an empty diff, so the added files are unknown - # and saying "none" here would be a clean bill of health nobody earned. - # Say why here: the caller only learns that nothing could be resolved, - # and "no version control is active" would be the wrong diagnosis when - # git is active and it is the range that failed. - print( - f'git is active but {_GIT_HEAD_RANGE} does not resolve, so what' - ' this change adds cannot be read from it. A shallow clone does' - ' this; fetch enough history for HEAD to have a parent.', - file=sys.stderr, - ) - return None - added = {f for f in head_diff.splitlines() if f.strip()} - return added | _git_renamed_to_new_names( - ['git', 'diff', _GIT_HEAD_RANGE], root - ) - - # 2. jj - if shutil.which('jj'): - code, jj_root = _run_cmd(['jj', 'root'], cwd=root) - if code == 0: - _, out = _run_cmd(['jj', 'diff', '--summary'], cwd=root) - added = set() - for line in out.splitlines(): - if line.startswith('A '): - parts = line.split(maxsplit=1) - if len(parts) == 2: - p = parts[1].strip() - if jj_root and not os.path.isabs(p): - p = os.path.join(jj_root, p) - added.add(p.replace(os.sep, '/')) - return added - - # 3. hg - if shutil.which('hg'): - code, hg_root = _run_cmd(['hg', 'root'], cwd=root) - if code == 0: - # A bare `hg status --added` reports only files added and not yet - # committed, so it goes empty the moment the change is committed or - # amended -- which is the usual state of a checkout by the time anyone - # runs this. Diff against the last synced revision instead, so the file - # set is the change's own content whether or not it is committed. - # `_SYNCED_BASE` degrades to '.' in a checkout with no local commits, - # where it reports the same thing a bare status does. - code, out = _run_cmd( - ['hg', 'status', '--added', '--no-status', '--rev', _SYNCED_BASE], - cwd=root, - ) - if code != 0: - # A repository whose phases do not distinguish local work this way. - _, out = _run_cmd(['hg', 'status', '--added', '--no-status'], cwd=root) - return { - ( - os.path.join(hg_root, f.strip()) - if (hg_root and not os.path.isabs(f.strip())) - else f.strip() - ).replace(os.sep, '/') - for f in out.splitlines() - if f.strip() - } - - # 4. g4 - if shutil.which('g4'): - code, _ = _run_cmd(['g4', 'info'], cwd=root) - if code == 0: - _, out = _run_cmd(['g4', 'opened'], cwd=root) - added = set() - for line in out.splitlines(): - if ' - add ' in line: - depot_file = line.split(' - add ')[0].split('#')[0].strip() - added.add(depot_file) - return added - - # 5. p4 - if shutil.which('p4'): - code, _ = _run_cmd(['p4', 'info'], cwd=root) - if code == 0: - _, out = _run_cmd(['p4', 'opened'], cwd=root) - added = set() - for line in out.splitlines(): - if ' - add ' in line: - depot_file = line.split(' - add ')[0].split('#')[0].strip() - added.add(depot_file) - return added - - return None + if not _in_git_work_tree(root): + return None + mid_commit = _git_is_mid_commit(root) + if mid_commit is None: + print( + 'This is a git work tree, but its index cannot be read, so whether' + ' this change is staged or committed is unknown, and so is what it' + ' adds.', + file=sys.stderr, + ) + return None + if mid_commit: + return _git_added_paths(['git', 'diff', '--cached'], root) + range_code, head_diff = _run_cmd( + ['git', 'diff', _GIT_HEAD_RANGE, '--name-only', _GIT_ADD_FILTER], + cwd=root, + ) + if range_code != 0: + # HEAD~1 is unreachable, as in a depth-1 clone. The range resolved to + # nothing rather than to an empty diff, so the added files are unknown and + # saying "none" here would be a clean bill of health nobody earned. Say why + # here: the caller only learns that nothing could be resolved, and would + # otherwise leave "not a git work tree" as the likeliest reading when it is + # the range that failed. + print( + f'This is a git work tree, but {_GIT_HEAD_RANGE} does not resolve, so' + ' what this change adds cannot be read from it. A shallow clone does' + ' this; fetch enough history for HEAD to have a parent.', + file=sys.stderr, + ) + return None + added = {f for f in head_diff.splitlines() if f.strip()} + return added | _git_renamed_to_new_names( + ['git', 'diff', _GIT_HEAD_RANGE], root + ) def get_commit_message(root: str = '.') -> str: - """Retrieves commit message or description from VCS. + """Retrieves the commit messages to search for a waiver, using git. Args: root: The root directory of the repository. Returns: - The message to search for a waiver tag, or '' when none can be read. + The messages to search for a waiver tag, or '' when none can be read. """ - # 1. git - if shutil.which('git'): - code, _ = _run_cmd(['git', 'rev-parse', '--is-inside-work-tree'], cwd=root) - if code == 0: - if _git_is_mid_commit(root) is not False: - # The change is staged, so the commit carrying it does not exist yet - # and its message is nowhere to be read: a pre-commit hook runs before - # git records what the author typed, and HEAD still describes the - # previous change. Returning HEAD's message here is what let a waiver - # written for an earlier commit silently cover this one. Waiving the - # change being committed goes through the environment instead -- - # `NO_UNIT_GUIDE='' git commit ...` -- which - # has_no_unit_guide_tag honours and the violation text advertises. - # None lands here too: a waiver that cannot be attributed to a change - # must not be applied to one. - return '' - _, msg = _run_cmd(['git', 'log', '-1', '--pretty=%B'], cwd=root) - # On a pull request, HEAD is a merge commit whose own message is - # generated by CI and can hold no waiver. The commits being merged are - # the ones the contributor wrote, so read the same range the added-file - # scan falls back to. Empty when HEAD~1 is unreachable. - _, range_msg = _run_cmd( - ['git', 'log', _GIT_HEAD_RANGE, '--pretty=%B'], cwd=root - ) - if range_msg: - msg = f'{msg}\n{range_msg}' - # COMMIT_EDITMSG is deliberately not consulted. It was read here to - # catch the message of the commit being made, which it never held: git - # writes it only after the pre-commit hook has run, so during that hook - # it carries the previous commit's message, or the message of an attempt - # some hook rejected. Both are messages written for another change, and - # neither can be told from a current one by inspection. - return msg - - # 2. jj - if shutil.which('jj'): - code, _ = _run_cmd(['jj', 'root'], cwd=root) - if code == 0: - _, out = _run_cmd( - ['jj', 'log', '-r', '@', '--no-graph', '-T', 'description'], cwd=root - ) - return out - - # 3. hg - if shutil.which('hg'): - code, _ = _run_cmd(['hg', 'root'], cwd=root) - if code == 0: - # Every local commit, not just the tip. The added-file scan above spans - # the whole range back to the last synced revision, so reading only the - # tip's message would let one commit on top bury a waiver written in the - # commit that actually adds the file -- and would let an unrelated tip - # message waive the whole range. - code, out = _run_cmd( - ['hg', 'log', '-r', _LOCAL_COMMITS, '--template', '{desc}\n'], - cwd=root, - ) - if code != 0: - _, out = _run_cmd( - ['hg', 'log', '-r', '.', '--template', '{desc}'], cwd=root - ) - return out - - # 4. g4 - if shutil.which('g4'): - code, _ = _run_cmd(['g4', 'info'], cwd=root) - if code == 0: - code, out = _run_cmd(['g4', 'change', '-o'], cwd=root) - if code == 0 and out: - return out - _, out = _run_cmd(['g4', 'describe'], cwd=root) - return out - - # 5. p4 - if shutil.which('p4'): - code, _ = _run_cmd(['p4', 'info'], cwd=root) - if code == 0: - _, out = _run_cmd(['p4', 'change', '-o'], cwd=root) - return out - - return '' + if not _in_git_work_tree(root): + return '' + if _git_is_mid_commit(root) is not False: + # The change is staged, so the commit carrying it does not exist yet and + # its message is nowhere to be read: a pre-commit hook runs before git + # records what the author typed, and HEAD still describes the previous + # change. Returning HEAD's message here is what let a waiver written for an + # earlier commit silently cover this one. Waiving the change being + # committed goes through the environment instead -- + # `NO_UNIT_GUIDE='' git commit ...` -- which has_no_unit_guide_tag + # honours and the violation text advertises. None lands here too: a waiver + # that cannot be attributed to a change must not be applied to one. + return '' + _, msg = _run_cmd(['git', 'log', '-1', '--pretty=%B'], cwd=root) + # On a pull request, HEAD is a merge commit whose own message is generated by + # CI and can hold no waiver. The commits being merged are the ones the + # contributor wrote, so read the same range the added-file scan falls back + # to. Empty when HEAD~1 is unreachable. + _, range_msg = _run_cmd( + ['git', 'log', _GIT_HEAD_RANGE, '--pretty=%B'], cwd=root + ) + if range_msg: + msg = f'{msg}\n{range_msg}' + # COMMIT_EDITMSG is deliberately not consulted. It was read here to catch the + # message of the commit being made, which it never held: git writes it only + # after the pre-commit hook has run, so during that hook it carries the + # previous commit's message, or the message of an attempt some hook rejected. + # Both are messages written for another change, and neither can be told from + # a current one by inspection. + return msg def is_exempt_from_unit_guide(rel_path: str, filename: str) -> bool: @@ -598,40 +475,15 @@ def has_no_unit_guide_tag(commit_msg: str) -> bool: return bool(_NO_UNIT_GUIDE_TAG.search(commit_msg)) -def _depot_path_to_abs(depot_path: str, adk_real_root: str) -> str: - """Locates a depot-style path inside the package being checked. - - A depot path names a file by its position in the repository the VCS serves, - which shares no prefix with the checkout on disk. What the two do share is - the package itself, so the split point is the last occurrence of the import - path. Keying on that rather than on a repository prefix keeps any particular - repository layout out of this script. - - Args: - depot_path: A path of the form `///<...>/.py`. - adk_real_root: Absolute, symlink-resolved path of the package root. - - Returns: - The absolute path of the matching file, or the depot path resolved as-is - when it does not run through the package. The caller drops anything that - does not land inside the package. - """ - clean_path = depot_path.lstrip('/') - marker = f'{_PACKAGE_IMPORT_PATH}/' - if marker in clean_path: - rel_to_package = clean_path.rsplit(marker, 1)[1] - return os.path.realpath(os.path.join(adk_real_root, rel_to_package)) - return os.path.realpath(clean_path) - - def _subpackage_renames(package_dir: str) -> dict[str, str]: """Maps a subpackage's real directory to the name the source tree gives it. - A subpackage can be exposed under a name of its own: `dependencies` points - at `dependencies_external`. Which of the two names a path arrives wearing - depends only on how it was detected -- git reports it relative to the - checkout, while the Piper-shaped detectors report the real location -- so - without this the same file demands its guide in two different directories. + A checkout can expose a subpackage through a symlink whose name differs from + the directory it points at. Which of the two names a path arrives wearing + depends only on how it was given -- a path relative to the checkout keeps the + link's name, while an absolute path into the linked directory, which only an + explicit file list supplies, carries the directory's -- so without this the + same file demands its guide in two different directories. Args: package_dir: The checkout's own `src/google/adk`, symlinks unresolved. @@ -708,7 +560,7 @@ def _normalize_and_filter_files( """Normalizes added files and filters to relevant Python source files. Handles both standard layout (src/google/adk/) and symlinked package - structures where subpackages point to an upstream source tree. + structures where subpackages link to a source tree kept elsewhere. Args: raw_files: The set of raw file paths to normalize and filter. @@ -738,44 +590,37 @@ def _normalize_and_filter_files( if not raw_file or not raw_file.endswith('.py'): continue - # Handle depot-style paths (e.g. //depot/.../agents/_agent.py), which the - # Perforce-style branches of get_vcs_added_files report. - if raw_file.startswith('//'): - abs_file = _depot_path_to_abs(raw_file, adk_real_root) - elif os.path.isabs(raw_file): + if os.path.isabs(raw_file): abs_file = os.path.realpath(raw_file) else: abs_file = os.path.realpath(os.path.join(repo_root, raw_file)) # Before resolving anything, see whether the path already sits under the - # checkout's own src/google/adk. A subpackage exposed there under a name - # different from its own -- `dependencies` for `dependencies_external` -- - # would otherwise resolve through the symlink and come back wearing the - # name the source tree does not use, so the guide would be demanded at a - # directory that does not exist in the exported repository. Keeping the - # unresolved form makes the source-tree name win, which is the one both - # this checkout and the exported repository agree on. + # checkout's own src/google/adk. A subpackage exposed there through a + # symlink named differently from its directory would otherwise resolve + # through the link and come back wearing the directory's name, so the + # guide would be demanded at a directory the source tree does not have. + # Keeping the unresolved form makes the source-tree name win. lexical = os.path.abspath(os.path.join(repo_root, raw_file)) - if not raw_file.startswith('//') and lexical.startswith( - package_dir + os.sep - ): + if lexical.startswith(package_dir + os.sep): rel_to_adk = os.path.relpath(lexical, package_dir).replace(os.sep, '/') if _keep_relative_path(rel_to_adk): results.append((raw_file, rel_to_adk, os.path.basename(lexical))) continue - # Check whether the file belongs to the package in either internal or - # external layout. + # Check whether the file belongs to the package, both when the checkout + # holds the package itself and when it exposes it through symlinks. # - # The checkout's own src/google/adk comes first, and must: internally it - # sits *inside* the package it points into, so both prefixes match a file - # under it and matching the outer one first mislabels the file. A path - # there resolves out to the package only through a subpackage symlink, so - # a file in a subpackage the checkout has no symlink for -- a subpackage - # the change is adding -- stays put and relativizes against the package - # root as `/src/google/adk/<...>`. That starts with an excluded - # directory name, so the file was dropped and a change adding a new - # subpackage passed both rules without being examined. + # The checkout's own src/google/adk comes first, and must: a checkout that + # exposes the package through symlinks can sit *inside* the package it + # points into, so both prefixes match a file under it and matching the + # outer one first mislabels the file. A path there resolves out to the + # package only through a subpackage symlink, so a file in a subpackage the + # checkout has no symlink for -- a subpackage the change is adding -- stays + # put and relativizes against the package root as + # `/src/google/adk/<...>`. That starts with an excluded directory + # name, so the file was dropped and a change adding a new subpackage passed + # both rules without being examined. if abs_file.startswith(package_real_dir + os.sep): rel_to_adk = os.path.relpath(abs_file, package_real_dir).replace( os.sep, '/' @@ -889,7 +734,7 @@ def _parse_args(argv: list[str]) -> argparse.Namespace: '--baseline-dir', help=( 'Baseline source tree to diff against (an origin/main checkout). If' - ' omitted, detects added files via local VCS.' + ' omitted, detects added files via git.' ), ) parser.add_argument( @@ -958,9 +803,9 @@ def main(argv: list[str]) -> int: ) return _EXIT_SETUP_ERROR - # Only a VCS can supply a commit message, so this is '' when checking an - # exported tree that has none. NO_UNIT_GUIDE comes from the environment - # there instead -- see _GUIDE_VIOLATION_LINE. + # Only git can supply a commit message, so this is '' when checking a tree + # that is not a git work tree. NO_UNIT_GUIDE comes from the environment there + # instead -- see _GUIDE_VIOLATION_LINE. commit_msg = get_commit_message(repo_root) if args.added_files_from: if not os.path.isfile(args.added_files_from): @@ -983,20 +828,41 @@ def main(argv: list[str]) -> int: return _EXIT_SETUP_ERROR raw_added_files = added_py_files_from_baseline(repo_root, args.baseline_dir) else: - vcs_added = get_vcs_added_files(repo_root) - if vcs_added is None: + git_added = get_git_added_files(repo_root) + if git_added is None: print( 'Could not determine the added files: no --baseline-dir or' - ' --added-files-from was given, and no version control in' - f' {os.path.abspath(repo_root)} could report them -- either none of' - ' git/jj/hg/g4/p4 is active there, or the one that is could not' - ' resolve what this change added (see above).\n' + ' --added-files-from was given, and git could not report them for' + f' {os.path.abspath(repo_root)}: either git is not installed, the' + ' directory is not in a git work tree, or git could not resolve what' + ' this change added (see any message above).\n' 'This is not a clean bill of health -- nothing was checked. Pass' ' --baseline-dir or --added-files-from to say what to check.', file=sys.stderr, ) return _EXIT_INDETERMINATE - raw_added_files = vcs_added + raw_added_files = git_added + + # A `//`-prefixed .py name with nothing on disk behind it is a depot-style + # path from a version control server, and this script has no way to place it + # in the package. Checking the rest of the list without it would pass a file + # nobody looked at, so refuse the list. A `//` name that does exist -- POSIX + # allows the doubled slash, and Windows writes UNC paths that way -- is an + # ordinary path and is checked. + depot_style_paths = sorted( + p + for p in raw_added_files + if p.startswith('//') and p.endswith('.py') and not os.path.exists(p) + ) + if depot_style_paths: + print( + 'Error: these look like depot-style paths rather than files in this' + ' checkout, so they cannot be checked: ' + + ', '.join(depot_style_paths) + + '\nPass the paths of the files in the checkout instead.', + file=sys.stderr, + ) + return _EXIT_SETUP_ERROR filtered_files = _normalize_and_filter_files(raw_added_files, repo_root) diff --git a/tests/unittests/scripts/test_check_new_py_files.py b/tests/unittests/scripts/test_check_new_py_files.py index 4163d5f5f0e..aef83c8061d 100644 --- a/tests/unittests/scripts/test_check_new_py_files.py +++ b/tests/unittests/scripts/test_check_new_py_files.py @@ -16,7 +16,6 @@ from __future__ import annotations -import ntpath import os import pathlib import shutil @@ -181,37 +180,6 @@ def test_no_waiver_ignores_a_tag_the_caller_did_not_mean( assert check_new_py_files.main(argv + ['--no-waiver']) == 1 -def test_get_commit_message_hg_reads_every_local_commit( - monkeypatch: pytest.MonkeyPatch, -) -> None: - """The message range has to match the range the added-file scan covers. - - The file set spans back to the last synced revision, so reading only the - tip's message would let a commit stacked on top bury a waiver written in the - commit that adds the file. - """ - - def fake_which(cmd: str) -> str | None: - return '/usr/bin/' + cmd if cmd == 'hg' else None - - def fake_run_cmd(cmd: list[str], cwd: str | None = None) -> tuple[int, str]: - if cmd == ['hg', 'root']: - return 0, '/workspace' - if check_new_py_files._LOCAL_COMMITS in cmd: - return 0, 'add a seam\nNO_UNIT_GUIDE=internal\n\nlater unrelated commit\n' - if '-r' in cmd and '.' in cmd: - return 0, 'later unrelated commit' - return 1, '' - - monkeypatch.setattr(check_new_py_files.shutil, 'which', fake_which) - monkeypatch.setattr(check_new_py_files, '_run_cmd', fake_run_cmd) - monkeypatch.delenv('NO_UNIT_GUIDE', raising=False) - monkeypatch.delenv('SKIP_UNIT_GUIDE', raising=False) - - msg = check_new_py_files.get_commit_message('.') - assert check_new_py_files.has_no_unit_guide_tag(msg) - - def test_check_files_prefix_violation(tmp_path: pathlib.Path) -> None: # Missing '_' prefix files = [('src/google/adk/agents/agent.py', 'agents/agent.py', 'agent.py')] @@ -451,7 +419,7 @@ def test_main_baseline_dir_env_tag_waives_without_a_commit_message( ) -> None: """NO_UNIT_GUIDE works, and is advertised, where there is no commit message. - Baseline mode can run against an exported tree with no VCS, where + Baseline mode can run against a tree with no git history, where `get_commit_message` returns '', so the commit-message tag cannot be the only remedy the violation text offers. """ @@ -530,15 +498,15 @@ def test_main_checks_a_file_in_a_subpackage_with_no_symlink_yet( ) -> None: """A change that adds a whole new subpackage must still be checked. - Internally the checkout sits inside the package it points into, and its - src/google/adk reaches the real subpackages through per-subpackage symlinks. + A checkout can sit inside the package it points into, with its + src/google/adk reaching the real subpackages through per-subpackage symlinks. A subpackage the change is adding has no symlink yet, so its files stay put and used to relativize against the package root as `/src/google/adk/...` -- a path whose first component is an excluded directory name, so it was dropped and the change passed without being examined. """ - # The internal layout: a package root that *contains* the checkout. + # A package root that *contains* the checkout. package_root = tmp_path / 'pkg' checkout = package_root / 'checkout' real_agents = package_root / 'agents' @@ -695,22 +663,21 @@ def test_sh_forwarder_execution(tmp_path: pathlib.Path) -> None: def test_symlinked_layout_normalization(tmp_path: pathlib.Path) -> None: - # Simulate symlinked layout where open_source_workspace/src/google/adk/__init__.py - # is a symlink pointing to the real upstream package root. - upstream_adk = tmp_path / 'repo' / 'third_party' / 'adk' - upstream_adk.mkdir(parents=True) - (upstream_adk / '__init__.py').write_text('', encoding='utf-8') - - oss_workspace = upstream_adk / 'open_source_workspace' - oss_src_adk = oss_workspace / 'src' / 'google' / 'adk' - oss_src_adk.mkdir(parents=True) - # Symlink __init__.py pointing back to upstream_adk/__init__.py - (oss_src_adk / '__init__.py').symlink_to(upstream_adk / '__init__.py') - - # A file added in upstream package tree - added_file = str(upstream_adk / 'agents' / '_agent.py') + # A checkout nested inside the package root, whose src/google/adk/__init__.py + # is a symlink to the package's own. + package_root = tmp_path / 'pkg' + package_root.mkdir(parents=True) + (package_root / '__init__.py').write_text('', encoding='utf-8') + + checkout = package_root / 'checkout' + checkout_adk = checkout / 'src' / 'google' / 'adk' + checkout_adk.mkdir(parents=True) + (checkout_adk / '__init__.py').symlink_to(package_root / '__init__.py') + + # A file added in the package itself. + added_file = str(package_root / 'agents' / '_agent.py') results = check_new_py_files._normalize_and_filter_files( - [added_file], repo_root=str(oss_workspace) + [added_file], repo_root=str(checkout) ) assert len(results) == 1 display_path, rel_to_adk, filename = results[0] @@ -718,7 +685,9 @@ def test_symlinked_layout_normalization(tmp_path: pathlib.Path) -> None: assert filename == '_agent.py' -def test_get_vcs_added_files_git(monkeypatch: pytest.MonkeyPatch) -> None: +def test_get_git_added_files_reads_staged_additions( + monkeypatch: pytest.MonkeyPatch, +) -> None: def fake_which(cmd: str) -> str | None: return '/usr/bin/' + cmd if cmd == 'git' else None @@ -732,11 +701,11 @@ def fake_run_cmd(cmd: list[str], cwd: str | None = None) -> tuple[int, str]: monkeypatch.setattr(check_new_py_files.shutil, 'which', fake_which) monkeypatch.setattr(check_new_py_files, '_run_cmd', fake_run_cmd) - added = check_new_py_files.get_vcs_added_files('.') + added = check_new_py_files.get_git_added_files('.') assert added == {'src/google/adk/agents/_staged.py'} -def test_get_vcs_added_files_git_head_diff( +def test_get_git_added_files_reads_the_last_commit_when_nothing_is_staged( monkeypatch: pytest.MonkeyPatch, ) -> None: def fake_which(cmd: str) -> str | None: @@ -754,11 +723,11 @@ def fake_run_cmd(cmd: list[str], cwd: str | None = None) -> tuple[int, str]: monkeypatch.setattr(check_new_py_files.shutil, 'which', fake_which) monkeypatch.setattr(check_new_py_files, '_run_cmd', fake_run_cmd) - added = check_new_py_files.get_vcs_added_files('.') + added = check_new_py_files.get_git_added_files('.') assert added == {'src/google/adk/agents/_committed.py'} -def test_get_vcs_added_files_git_unreachable_range_is_indeterminate( +def test_get_git_added_files_unreachable_range_is_indeterminate( monkeypatch: pytest.MonkeyPatch, ) -> None: """A range that does not resolve is unknown, not empty. @@ -783,10 +752,10 @@ def fake_run_cmd(cmd: list[str], cwd: str | None = None) -> tuple[int, str]: monkeypatch.setattr(check_new_py_files.shutil, 'which', fake_which) monkeypatch.setattr(check_new_py_files, '_run_cmd', fake_run_cmd) - assert check_new_py_files.get_vcs_added_files('.') is None + assert check_new_py_files.get_git_added_files('.') is None -def test_get_vcs_added_files_git_empty_range_is_no_files( +def test_get_git_added_files_empty_range_is_no_files( monkeypatch: pytest.MonkeyPatch, ) -> None: """A range that resolves to an empty diff really is no added files.""" @@ -802,189 +771,96 @@ def fake_run_cmd(cmd: list[str], cwd: str | None = None) -> tuple[int, str]: monkeypatch.setattr(check_new_py_files.shutil, 'which', fake_which) monkeypatch.setattr(check_new_py_files, '_run_cmd', fake_run_cmd) - assert check_new_py_files.get_vcs_added_files('.') == set() - + assert check_new_py_files.get_git_added_files('.') == set() -def _patch_windows_paths(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setattr(check_new_py_files.os, 'path', ntpath) - monkeypatch.setattr(check_new_py_files.os, 'sep', '\\') - -@pytest.mark.parametrize('windows', [False, True]) -def test_get_vcs_added_files_jj( - monkeypatch: pytest.MonkeyPatch, windows: bool +@pytest.mark.parametrize( + 'git_installed', [True, False], ids=['git_installed', 'git_missing'] +) +def test_added_files_and_message_are_unknown_outside_a_git_work_tree( + monkeypatch: pytest.MonkeyPatch, git_installed: bool ) -> None: - def fake_which(cmd: str) -> str | None: - return '/usr/bin/' + cmd if cmd == 'jj' else None - - def fake_run_cmd(cmd: list[str], cwd: str | None = None) -> tuple[int, str]: - if cmd == ['jj', 'root']: - return 0, r'C:\workspace' if windows else '/workspace' - if cmd == ['jj', 'diff', '--summary']: - return 0, 'A src/google/adk/agents/_jj_agent.py\nM existing.py' - return 1, '' - - monkeypatch.setattr(check_new_py_files.shutil, 'which', fake_which) - monkeypatch.setattr(check_new_py_files, '_run_cmd', fake_run_cmd) - if windows: - _patch_windows_paths(monkeypatch) - - expected = ( - 'C:/workspace/src/google/adk/agents/_jj_agent.py' - if windows - else '/workspace/src/google/adk/agents/_jj_agent.py' - ) - added = check_new_py_files.get_vcs_added_files('.') - assert added == {expected} - + """Outside a git work tree the added files and the message are unknown. -_HG_SYNCED_BASE_STATUS = [ - 'hg', - 'status', - '--added', - '--no-status', - '--rev', - check_new_py_files._SYNCED_BASE, -] -_HG_WORKING_DIR_STATUS = ['hg', 'status', '--added', '--no-status'] + git's own answer about the work tree decides that. Another tool on PATH may + well answer, but the check reports that it could not tell rather than taking + another system's word for what the change is, whether git is merely not + managing root or not installed at all. + """ + consulted: list[str] = [] + def fake_run_cmd(cmd: list[str], cwd: str | None = None) -> tuple[int, str]: + consulted.append(cmd[0]) + if cmd[0] != 'git': + return 0, 'an answer' + # Only the work-tree probe fails, so the test pins that it is the probe, + # not some later git command, that ends the search. + return (1, '') if 'rev-parse' in cmd else (0, '') -@pytest.mark.parametrize('windows', [False, True]) -def test_get_vcs_added_files_hg( - monkeypatch: pytest.MonkeyPatch, windows: bool -) -> None: def fake_which(cmd: str) -> str | None: - return '/usr/bin/' + cmd if cmd == 'hg' else None - - def fake_run_cmd(cmd: list[str], cwd: str | None = None) -> tuple[int, str]: - if cmd == ['hg', 'root']: - return 0, r'C:\workspace' if windows else '/workspace' - if cmd == _HG_SYNCED_BASE_STATUS: - return 0, 'src/google/adk/agents/_hg_agent.py' - return 1, '' + if cmd == 'git' and not git_installed: + return None + return '/usr/bin/' + cmd monkeypatch.setattr(check_new_py_files.shutil, 'which', fake_which) monkeypatch.setattr(check_new_py_files, '_run_cmd', fake_run_cmd) - if windows: - _patch_windows_paths(monkeypatch) - expected = ( - 'C:/workspace/src/google/adk/agents/_hg_agent.py' - if windows - else '/workspace/src/google/adk/agents/_hg_agent.py' - ) - added = check_new_py_files.get_vcs_added_files('.') - assert added == {expected} + assert check_new_py_files.get_git_added_files('.') is None + assert check_new_py_files.get_commit_message('.') == '' + assert set(consulted) <= {'git'} -def test_get_vcs_added_files_hg_sees_an_already_committed_add( - monkeypatch: pytest.MonkeyPatch, +@pytest.mark.parametrize('channel', ['list_file', 'argument']) +def test_main_refuses_a_depot_style_name_that_is_not_a_file( + tmp_path: pathlib.Path, + capsys: pytest.CaptureFixture[str], + channel: str, ) -> None: - """A Mercurial checkout is normally committed by the time this runs. + """A `//`-prefixed .py name with no file behind it fails the run. - `hg status --added` on its own reports only files added and not yet - committed, so it goes empty after `hg commit` or `hg amend` and the check - silently passed every such change. The file set has to come from a diff - against the last synced revision instead. + Nothing can place it in the package, so without the refusal it would drop + out of the check without a word and the compliant file beside it would pass + alone. The refusal holds however the name is given. """ + new_dir = tmp_path / 'new' + compliant = _tree_with_added_file(new_dir, 'agents/_compliant.py') + depot_style = '//server/src/google/adk/agents/public.py' + argv = ['--new-dir', str(new_dir), '--no-unit-guide'] + if channel == 'list_file': + listing = tmp_path / 'added.txt' + listing.write_text(f'{compliant}\n{depot_style}\n', encoding='utf-8') + argv += ['--added-files-from', str(listing)] + else: + argv += [str(compliant), depot_style] - def fake_which(cmd: str) -> str | None: - return '/usr/bin/' + cmd if cmd == 'hg' else None - - def fake_run_cmd(cmd: list[str], cwd: str | None = None) -> tuple[int, str]: - if cmd == ['hg', 'root']: - return 0, '/workspace' - if cmd == _HG_SYNCED_BASE_STATUS: - return 0, 'src/google/adk/agents/_committed.py' - if cmd == _HG_WORKING_DIR_STATUS: - return 0, '' # Committed, so nothing is pending in the working copy. - return 1, '' - - monkeypatch.setattr(check_new_py_files.shutil, 'which', fake_which) - monkeypatch.setattr(check_new_py_files, '_run_cmd', fake_run_cmd) - - added = check_new_py_files.get_vcs_added_files('.') - assert added == {'/workspace/src/google/adk/agents/_committed.py'} + assert check_new_py_files.main(argv) == check_new_py_files._EXIT_SETUP_ERROR + assert depot_style in capsys.readouterr().err -def test_get_vcs_added_files_hg_falls_back_when_the_revset_fails( - monkeypatch: pytest.MonkeyPatch, +def test_main_checks_a_real_path_that_starts_with_a_double_slash( + tmp_path: pathlib.Path, ) -> None: - """A plain hg repository need not have the phases the revset relies on.""" + """A real file named with a leading `//` is checked, not refused. - def fake_which(cmd: str) -> str | None: - return '/usr/bin/' + cmd if cmd == 'hg' else None - - def fake_run_cmd(cmd: list[str], cwd: str | None = None) -> tuple[int, str]: - if cmd == ['hg', 'root']: - return 0, '/workspace' - if cmd == _HG_SYNCED_BASE_STATUS: - return 255, '' - if cmd == _HG_WORKING_DIR_STATUS: - return 0, 'src/google/adk/agents/_hg_agent.py' - return 1, '' - - monkeypatch.setattr(check_new_py_files.shutil, 'which', fake_which) - monkeypatch.setattr(check_new_py_files, '_run_cmd', fake_run_cmd) - - added = check_new_py_files.get_vcs_added_files('.') - assert added == {'/workspace/src/google/adk/agents/_hg_agent.py'} - - -def test_get_vcs_added_files_g4(monkeypatch: pytest.MonkeyPatch) -> None: - def fake_which(cmd: str) -> str | None: - return '/usr/bin/' + cmd if cmd == 'g4' else None - - def fake_run_cmd(cmd: list[str], cwd: str | None = None) -> tuple[int, str]: - if cmd == ['g4', 'info']: - return 0, 'Server: ...' - if cmd == ['g4', 'opened']: - return ( - 0, - ( - '//depot/mirror/src/google/adk/agents/_g4_agent.py#1' - ' - add default change (text)' - ), - ) - return 1, '' - - monkeypatch.setattr(check_new_py_files.shutil, 'which', fake_which) - monkeypatch.setattr(check_new_py_files, '_run_cmd', fake_run_cmd) - - added = check_new_py_files.get_vcs_added_files('.') - assert added == {'//depot/mirror/src/google/adk/agents/_g4_agent.py'} - - -def test_get_vcs_added_files_p4(monkeypatch: pytest.MonkeyPatch) -> None: - def fake_which(cmd: str) -> str | None: - return '/usr/bin/' + cmd if cmd == 'p4' else None - - def fake_run_cmd(cmd: list[str], cwd: str | None = None) -> tuple[int, str]: - if cmd == ['p4', 'info']: - return 0, 'Server: ...' - if cmd == ['p4', 'opened']: - return ( - 0, - ( - '//depot/mirror/src/google/adk/agents/_p4_agent.py#1' - ' - add default change (text)' - ), - ) - return 1, '' - - monkeypatch.setattr(check_new_py_files.shutil, 'which', fake_which) - monkeypatch.setattr(check_new_py_files, '_run_cmd', fake_run_cmd) - - added = check_new_py_files.get_vcs_added_files('.') - assert added == {'//depot/mirror/src/google/adk/agents/_p4_agent.py'} + POSIX allows a doubled leading slash, and Windows writes UNC paths that way, + so such a name can point at a genuine file. A `//` entry that is not a .py + file is ignored like any other non-Python entry. + """ + new_dir = tmp_path / 'new' + public = _tree_with_added_file(new_dir, 'agents/public.py') + listing = tmp_path / 'added.txt' + listing.write_text(f'/{public}\n//server/BUILD\n', encoding='utf-8') + exit_code = check_new_py_files.main([ + '--new-dir', + str(new_dir), + '--added-files-from', + str(listing), + '--no-unit-guide', + ]) -def test_get_vcs_added_files_none_detected( - monkeypatch: pytest.MonkeyPatch, -) -> None: - monkeypatch.setattr(check_new_py_files.shutil, 'which', lambda _: None) - added = check_new_py_files.get_vcs_added_files('.') - assert added is None + # Checked, and the public name breaks the prefix rule. + assert exit_code == check_new_py_files._EXIT_VIOLATIONS def test_get_commit_message_git( @@ -1078,108 +954,7 @@ def fake_run_cmd(cmd: list[str], cwd: str | None = None) -> tuple[int, str]: assert check_new_py_files.has_no_unit_guide_tag(msg) -def test_get_commit_message_jj( - monkeypatch: pytest.MonkeyPatch, tmp_path: pathlib.Path -) -> None: - def fake_which(cmd: str) -> str | None: - return '/usr/bin/' + cmd if cmd == 'jj' else None - - def fake_run_cmd(cmd: list[str], cwd: str | None = None) -> tuple[int, str]: - if cmd == ['jj', 'root']: - return 0, str(tmp_path) - if 'jj' in cmd and 'log' in cmd: - return 0, 'JJ Description' - return 1, '' - - monkeypatch.setattr(check_new_py_files.shutil, 'which', fake_which) - monkeypatch.setattr(check_new_py_files, '_run_cmd', fake_run_cmd) - - msg = check_new_py_files.get_commit_message(str(tmp_path)) - assert msg == 'JJ Description' - - -def test_get_commit_message_hg( - monkeypatch: pytest.MonkeyPatch, tmp_path: pathlib.Path -) -> None: - def fake_which(cmd: str) -> str | None: - return '/usr/bin/' + cmd if cmd == 'hg' else None - - def fake_run_cmd(cmd: list[str], cwd: str | None = None) -> tuple[int, str]: - if cmd == ['hg', 'root']: - return 0, str(tmp_path) - if 'hg' in cmd and 'log' in cmd: - return 0, 'HG Description' - return 1, '' - - monkeypatch.setattr(check_new_py_files.shutil, 'which', fake_which) - monkeypatch.setattr(check_new_py_files, '_run_cmd', fake_run_cmd) - - msg = check_new_py_files.get_commit_message(str(tmp_path)) - assert msg == 'HG Description' - - -def test_get_commit_message_g4( - monkeypatch: pytest.MonkeyPatch, tmp_path: pathlib.Path -) -> None: - def fake_which(cmd: str) -> str | None: - return '/usr/bin/' + cmd if cmd == 'g4' else None - - def fake_run_cmd(cmd: list[str], cwd: str | None = None) -> tuple[int, str]: - if cmd == ['g4', 'info']: - return 0, 'Server: ...' - if cmd == ['g4', 'change', '-o']: - return 0, 'G4 Change Description' - return 1, '' - - monkeypatch.setattr(check_new_py_files.shutil, 'which', fake_which) - monkeypatch.setattr(check_new_py_files, '_run_cmd', fake_run_cmd) - - msg = check_new_py_files.get_commit_message(str(tmp_path)) - assert msg == 'G4 Change Description' - - -def test_get_commit_message_p4( - monkeypatch: pytest.MonkeyPatch, tmp_path: pathlib.Path -) -> None: - def fake_which(cmd: str) -> str | None: - return '/usr/bin/' + cmd if cmd == 'p4' else None - - def fake_run_cmd(cmd: list[str], cwd: str | None = None) -> tuple[int, str]: - if cmd == ['p4', 'info']: - return 0, 'Server: ...' - if cmd == ['p4', 'change', '-o']: - return 0, 'P4 Change Description' - return 1, '' - - monkeypatch.setattr(check_new_py_files.shutil, 'which', fake_which) - monkeypatch.setattr(check_new_py_files, '_run_cmd', fake_run_cmd) - - msg = check_new_py_files.get_commit_message(str(tmp_path)) - assert msg == 'P4 Change Description' - - -def test_normalize_depot_path(tmp_path: pathlib.Path) -> None: - upstream_adk = tmp_path / 'third_party' / 'py' / 'google' / 'adk' - upstream_adk.mkdir(parents=True) - (upstream_adk / '__init__.py').write_text('', encoding='utf-8') - - workspace = upstream_adk / 'open_source_workspace' - src_adk = workspace / 'src' / 'google' / 'adk' - src_adk.mkdir(parents=True) - (src_adk / '__init__.py').symlink_to(upstream_adk / '__init__.py') - - depot_path = '//depot/mirror/src/google/adk/agents/_g4_agent.py' - results = check_new_py_files._normalize_and_filter_files( - [depot_path], repo_root=str(workspace) - ) - assert len(results) == 1 - display_path, rel_to_adk, filename = results[0] - assert display_path == depot_path - assert rel_to_adk == 'agents/_g4_agent.py' - assert filename == '_g4_agent.py' - - -def test_main_no_vcs_no_baseline( +def test_main_without_git_or_baseline_is_indeterminate( tmp_path: pathlib.Path, monkeypatch: pytest.MonkeyPatch, capsys: pytest.CaptureFixture[str], @@ -1194,7 +969,7 @@ def test_main_no_vcs_no_baseline( exit_code = check_new_py_files.main(['--new-dir', str(new_dir)]) # 3, not 1 or 2: nothing was checked, which is neither a pass nor a - # violation. run_precommit_checks reports this as skipped. + # violation. A caller running this opportunistically reports it as skipped. assert exit_code == check_new_py_files._EXIT_INDETERMINATE err = capsys.readouterr().err assert 'Could not determine the added files' in err @@ -1220,8 +995,8 @@ def test_sh_forwarder_execution_from_any_cwd(tmp_path: pathlib.Path) -> None: # Tests that drive a real repository rather than monkeypatching _run_cmd. The -# faked tests above pin the parsing of each VCS's output; these pin what the -# VCS actually says, which is where the interesting mistakes live -- a rename +# faked tests above pin the parsing of git's output; these pin what git +# actually says, which is where the interesting mistakes live -- a rename # reported as R100 rather than as an add, for one. @@ -1236,37 +1011,33 @@ def test_a_renamed_subpackage_keeps_its_source_tree_name( ) -> None: """A subpackage exposed under another name is checked under that name. - `dependencies` points at `dependencies_external`. Resolving the symlink - would report the target's name, and the guide would then be demanded at a - directory that does not exist in the tree the contributor sees. + `dependencies` is a symlink to a directory named differently. Resolving the + symlink would report the target's name, and the guide would then be demanded + at a directory that does not exist in the tree the contributor sees. """ - # The internal shape: a package root holding the real subpackage, and a - # checkout inside it whose src/google/adk exposes it under another name. + # A package root holding the real subpackage, and a checkout inside it whose + # src/google/adk exposes it under another name. package_root = tmp_path / 'pkg' - (package_root / 'dependencies_external').mkdir(parents=True) + (package_root / 'dependencies_impl').mkdir(parents=True) (package_root / '__init__.py').write_text('', encoding='utf-8') - added = package_root / 'dependencies_external' / '_thing.py' + added = package_root / 'dependencies_impl' / '_thing.py' added.write_text('', encoding='utf-8') checkout = package_root / 'checkout' adk_src = checkout / 'src' / 'google' / 'adk' adk_src.mkdir(parents=True) (checkout / 'docs' / 'guides').mkdir(parents=True) - os.symlink(package_root / 'dependencies_external', adk_src / 'dependencies') + os.symlink(package_root / 'dependencies_impl', adk_src / 'dependencies') os.symlink(package_root / '__init__.py', adk_src / '__init__.py') # Which of the subpackage's two names a path arrives wearing depends only on - # how it was detected: git reports it relative to the checkout, while the - # Piper-shaped detectors report the real location. All of them have to land - # on the name the source tree uses, or the same file demands its guide in - # two different directories depending on where it is checked. + # how it was given: relative to the checkout, as git and --baseline-dir give + # it, or as an absolute path into the real subpackage. Both have to land on + # the name the source tree uses, or the same file demands its guide in two + # different directories depending on how it was named. for raw in ( - 'src/google/adk/dependencies/_thing.py', # git, --baseline-dir - str(added), # hg and jj, an absolute path into the real subpackage - ( # g4 and p4 - '//depot/mirror/third_party/py/google/adk/' - 'dependencies_external/_thing.py' - ), + 'src/google/adk/dependencies/_thing.py', # relative, through the link + str(added), # absolute, into the linked directory ): results = check_new_py_files._normalize_and_filter_files( [raw], repo_root=str(checkout) @@ -1336,7 +1107,7 @@ def test_real_git_flags_a_rename_into_a_public_name( ) _git(repo, 'commit', '-qm', 'refactor: rename') - added = check_new_py_files.get_vcs_added_files(str(repo)) + added = check_new_py_files.get_git_added_files(str(repo)) assert added == {'src/google/adk/agents/brand_new_public.py'} @@ -1356,7 +1127,7 @@ def test_real_git_ignores_a_pure_relocation(tmp_path: pathlib.Path) -> None: ) _git(repo, 'commit', '-qm', 'refactor: relocate') - assert check_new_py_files.get_vcs_added_files(str(repo)) == set() + assert check_new_py_files.get_git_added_files(str(repo)) == set() def test_real_git_ignores_a_public_to_public_rename( @@ -1385,7 +1156,7 @@ def test_real_git_ignores_a_public_to_public_rename( ) _git(repo, 'commit', '-qm', 'refactor: rename') - assert check_new_py_files.get_vcs_added_files(str(repo)) == set() + assert check_new_py_files.get_git_added_files(str(repo)) == set() def test_real_git_flags_a_file_moved_in_from_an_excluded_tree( @@ -1412,7 +1183,7 @@ def test_real_git_flags_a_file_moved_in_from_an_excluded_tree( ) _git(repo, 'commit', '-qm', 'promote the helper') - assert check_new_py_files.get_vcs_added_files(str(repo)) == { + assert check_new_py_files.get_git_added_files(str(repo)) == { 'src/google/adk/agents/helper_public.py' } @@ -1436,7 +1207,7 @@ def test_real_git_flags_a_stub_promoted_to_a_module( ) _git(repo, 'commit', '-qm', 'promote the stub') - assert check_new_py_files.get_vcs_added_files(str(repo)) == { + assert check_new_py_files.get_git_added_files(str(repo)) == { 'src/google/adk/agents/thing.py' } @@ -1466,7 +1237,7 @@ def test_real_git_flags_a_move_out_of_a_guide_exempt_subtree( ) _git(repo, 'commit', '-qm', 'move it out of cli') - assert check_new_py_files.get_vcs_added_files(str(repo)) == { + assert check_new_py_files.get_git_added_files(str(repo)) == { 'src/google/adk/agents/tool.py' } @@ -1507,7 +1278,22 @@ def test_real_git_staged_edit_is_not_judged_on_the_previous_commit( existing.write_text('# edited\n', encoding='utf-8') _git(repo, 'add', str(existing)) - assert check_new_py_files.get_vcs_added_files(str(repo)) == set() + assert check_new_py_files.get_git_added_files(str(repo)) == set() + + +def test_real_git_the_git_directory_is_not_a_work_tree( + tmp_path: pathlib.Path, capsys: pytest.CaptureFixture[str] +) -> None: + """Inside .git itself there is no work tree, and the check does not claim one. + + git answers that question with `false` and a zero exit status, so a probe + that trusted the status went on to report a work tree whose index could not + be read -- a diagnosis that sends the reader to the wrong problem. + """ + repo = _git_repo_with_a_guided_module(tmp_path) + + assert check_new_py_files.get_git_added_files(str(repo / '.git')) is None + assert 'git work tree, but' not in capsys.readouterr().err def test_real_git_reports_a_staged_addition(tmp_path: pathlib.Path) -> None: @@ -1517,7 +1303,7 @@ def test_real_git_reports_a_staged_addition(tmp_path: pathlib.Path) -> None: ) _git(repo, 'add', '-A') - assert check_new_py_files.get_vcs_added_files(str(repo)) == { + assert check_new_py_files.get_git_added_files(str(repo)) == { 'src/google/adk/agents/_added.py' } From b4c5272c53eb5401480ad7472535d1099ef4e9d7 Mon Sep 17 00:00:00 2001 From: Shangjie Chen Date: Tue, 29 Sep 2026 11:00:02 -0700 Subject: [PATCH 12/29] fix(tools): surface NodeTool failures to on_tool_error and return dict validation errors NodeTool swallowed every exception raised while running its node and returned a plain string, which the flow wrapped as {'result': ''}. As a result: - on_tool_error callbacks (agent and plugin) never saw node failures, unlike every other BaseTool, and the original exception was lost behind the generic "Dynamic node failed" wrapper. - Input validation errors reached the model as {'result': ...} instead of the {'error': ...} shape FunctionTool uses for argument validation errors. This change aligns NodeTool with FunctionTool: - Input schema validation errors return {'error': ...} with the same wording as FunctionTool, so the model can correct its arguments and retry. - When the node fails, the node's original exception is re-raised (chained from DynamicNodeFailError). The tool pipeline then runs on_tool_error callbacks with the real cause; if none handles it, the error propagates, as it does for FunctionTool. Node-level retry_config still applies first, inside run_node. Behavior change: an agent using a node as a tool without an on_tool_error callback now fails the run when the node fails, instead of passing an error string to the model. This matches FunctionTool. Co-authored-by: Shangjie Chen PiperOrigin-RevId: 990410996 --- src/google/adk/tools/_node_tool.py | 58 ++--- src/google/adk/workflow/_function_node.py | 67 ++++++ tests/unittests/agents/test_context.py | 2 + tests/unittests/workflow/test_node_tool.py | 241 +++++++++++++++++++++ 4 files changed, 341 insertions(+), 27 deletions(-) diff --git a/src/google/adk/tools/_node_tool.py b/src/google/adk/tools/_node_tool.py index 37cdec335eb..ed56daaa923 100644 --- a/src/google/adk/tools/_node_tool.py +++ b/src/google/adk/tools/_node_tool.py @@ -17,11 +17,13 @@ from typing import Any from google.genai import types +from pydantic import ValidationError from typing_extensions import override from ..utils._schema_utils import schema_to_json_schema from ..workflow._base_node import BaseNode -from ..workflow._errors import NodeInterruptedError +from ..workflow._errors import DynamicNodeFailError +from ..workflow._errors import WorkflowDataError from .base_tool import BaseTool from .tool_context import ToolContext @@ -131,27 +133,29 @@ async def run_async( args: dict[str, Any], tool_context: ToolContext, ) -> Any: - import inspect - - from pydantic import BaseModel - input_schema = getattr(self.node, 'input_schema', None) - node_input: Any - if inspect.isclass(input_schema) and issubclass(input_schema, BaseModel): - try: - node_input = input_schema.model_validate(args) - except Exception as e: - return f'Error validating input for node: {e}' + schema = ( + schema_to_json_schema(input_schema) + if input_schema is not None + else None + ) + if isinstance(schema, dict) and schema.get('type') != 'object': + node_input = args.get('request') else: - schema = ( - schema_to_json_schema(input_schema) - if input_schema is not None - else None - ) - if isinstance(schema, dict) and schema.get('type') != 'object': - node_input = args.get('request') - else: - node_input = args + node_input = args + + try: + node_input = self.node._validate_input_data(node_input) + except (ValidationError, WorkflowDataError) as e: + # Same shape as FunctionTool's argument validation errors, so the + # model can correct its arguments and retry. + return { + 'error': ( + f'Invoking `{self.name}()` failed due to argument validation' + f' errors:\n{e}\nYou could retry calling this tool with' + ' corrected argument types.' + ) + } fc_id = tool_context.function_call_id base_branch = tool_context.branch @@ -166,10 +170,10 @@ async def run_async( use_sub_branch=False, raise_on_wait=True, ) - if res is None: - return {'result': None} - return res - except NodeInterruptedError: - raise - except Exception as e: - return f'Error running node {self.name}: {e}' + except DynamicNodeFailError as e: + # Surface the node's own error, as a FunctionTool would, so the tool + # pipeline runs on_tool_error callbacks with the real cause. + raise e.error from e + if res is None: + return {'result': None} + return res diff --git a/src/google/adk/workflow/_function_node.py b/src/google/adk/workflow/_function_node.py index e7c589f576f..3e1291e17e5 100644 --- a/src/google/adk/workflow/_function_node.py +++ b/src/google/adk/workflow/_function_node.py @@ -338,6 +338,24 @@ def _bind_parameters(self, ctx: Context, node_input: Any) -> dict[str, Any]: except (TypeError, KeyError): pass + if ( + not has_param + and input_bound + and param_name == "node_input" + and self.input_schema is not None + and not isinstance(self.input_schema, (dict, types.Schema)) + and param_name in self._type_hints + ): + try: + value = self._coerce_param( + param_name, + node_input, + self._type_hints[param_name], + ) + has_param = True + except Exception: + pass + if has_param: if param_name in self._type_hints: value = self._coerce_param( @@ -437,6 +455,55 @@ def _coerce_param( adapter = TypeAdapter(annotated_type) return adapter.validate_python(value) + @override + def _validate_input_data(self, data: Any) -> Any: + """Validates input data for FunctionNode.""" + if self.input_schema is not None and not isinstance( + self.input_schema, (dict, types.Schema) + ): + return super()._validate_input_data(data) + + if self.parameter_binding == "node_input": + source: Any = data if isinstance(data, (dict, BaseModel)) else {} + validated: dict[str, Any] = {} + for param_name, param in self._sig.parameters.items(): + if param_name == self._context_param_name: + continue + + has_param = False + value = None + if isinstance(source, BaseModel): + if hasattr(source, param_name): + has_param = True + value = getattr(source, param_name) + else: + try: + if param_name in source: + has_param = True + value = source[param_name] + except (TypeError, KeyError): + pass + + if has_param: + if param_name in self._type_hints: + value = self._coerce_param( + param_name, + value, + self._type_hints[param_name], + ) + validated[param_name] = value + elif param.default is not inspect.Parameter.empty: + validated[param_name] = param.default + else: + raise WorkflowDataError( + f'Missing value for parameter "{param_name}" of function' + f' "{self.name}". It was not found in node_input and has no' + " default value." + ) + return validated + + return super()._validate_input_data(data) + @override def model_copy( self, *, update: Mapping[str, Any] | None = None, deep: bool = False diff --git a/tests/unittests/agents/test_context.py b/tests/unittests/agents/test_context.py index 8e6a00adf37..d7b1cf4f676 100644 --- a/tests/unittests/agents/test_context.py +++ b/tests/unittests/agents/test_context.py @@ -788,6 +788,8 @@ async def test_tool_context_from_node_ic_nests_agent_tool_and_node_tool_children mock_copy.branch = None mock_copy.invocation_id = "inv-1" mock_copy.session = mock_invocation_context.session + mock_copy._enqueue_event = AsyncMock() + mock_copy.model_copy.return_value = mock_copy mock_invocation_context.model_copy.return_value = mock_copy node_ic = caller_ctx.get_invocation_context() diff --git a/tests/unittests/workflow/test_node_tool.py b/tests/unittests/workflow/test_node_tool.py index 2d5135fab55..f16a3da19ea 100644 --- a/tests/unittests/workflow/test_node_tool.py +++ b/tests/unittests/workflow/test_node_tool.py @@ -1585,3 +1585,244 @@ def add(a: int, b: int) -> Any: if part.text ] assert texts == ['done'] + + +def _function_responses(events: list[Event]) -> list[Any]: + return [ + part.function_response.response + for event in events + if event.content and event.content.parts + for part in event.content.parts + if part.function_response + ] + + +class _TypedInput(BaseModel): + x: int + + +@pytest.mark.asyncio +async def test_node_tool_validation_error_returns_error_dict( + request: pytest.FixtureRequest, +): + """Invalid LLM args yield an {'error': ...} response, like FunctionTool.""" + + def typed(node_input: _TypedInput) -> int: + return node_input.x + + typed_node = FunctionNode(func=typed) + typed_node.input_schema = _TypedInput + agent = LlmAgent( + name='agent', + model=testing_utils.MockModel.create( + responses=[ + types.Part.from_function_call(name='typed', args={'x': 'abc'}), + types.Part.from_text(text='done'), + ] + ), + tools=[typed_node], + ) + runner = testing_utils.InMemoryRunner( + app=App(name=request.function.__name__, root_agent=agent) + ) + + events = await runner.run_async(testing_utils.get_user_content('go')) + + responses = _function_responses(events) + assert len(responses) == 1 + assert set(responses[0]) == {'error'} + assert 'argument validation errors' in responses[0]['error'] + + +@pytest.mark.asyncio +async def test_node_tool_function_node_validation_error_returns_error_dict( + request: pytest.FixtureRequest, +): + """FunctionNode tool returns {'error': ...} when arguments fail validation.""" + + def calculate(x: int) -> int: + return x + + calc_node = FunctionNode(func=calculate) + agent = LlmAgent( + name='agent', + model=testing_utils.MockModel.create( + responses=[ + types.Part.from_function_call( + name='calculate', args={'x': 'abc'} + ), + types.Part.from_text(text='done'), + ] + ), + tools=[calc_node], + ) + runner = testing_utils.InMemoryRunner( + app=App(name=request.function.__name__, root_agent=agent) + ) + + events = await runner.run_async(testing_utils.get_user_content('go')) + + responses = _function_responses(events) + assert len(responses) == 1 + assert set(responses[0]) == {'error'} + assert 'argument validation errors' in responses[0]['error'] + + +@pytest.mark.asyncio +async def test_node_tool_failure_reaches_on_tool_error_callback( + request: pytest.FixtureRequest, +): + """A failing node surfaces its original error to on_tool_error callbacks.""" + seen_errors: list[Exception] = [] + + def boom(x: int) -> int: + raise ValueError(f'kaboom {x}') + + def on_tool_error(tool, args, tool_context, error): + seen_errors.append(error) + return {'handled': True} + + agent = LlmAgent( + name='agent', + model=testing_utils.MockModel.create( + responses=[ + types.Part.from_function_call(name='boom', args={'x': 1}), + types.Part.from_text(text='done'), + ] + ), + tools=[FunctionNode(func=boom)], + on_tool_error_callback=on_tool_error, + ) + runner = testing_utils.InMemoryRunner( + app=App(name=request.function.__name__, root_agent=agent) + ) + + events = await runner.run_async(testing_utils.get_user_content('go')) + + assert len(seen_errors) == 1 + assert isinstance(seen_errors[0], ValueError) + assert str(seen_errors[0]) == 'kaboom 1' + assert _function_responses(events) == [{'handled': True}] + + +@pytest.mark.asyncio +async def test_node_tool_unhandled_failure_propagates( + request: pytest.FixtureRequest, +): + """Without on_tool_error, a node failure propagates like a FunctionTool's.""" + + def boom(x: int) -> int: + raise ValueError(f'kaboom {x}') + + agent = LlmAgent( + name='agent', + model=testing_utils.MockModel.create( + responses=[ + types.Part.from_function_call(name='boom', args={'x': 1}), + types.Part.from_text(text='done'), + ] + ), + tools=[FunctionNode(func=boom)], + ) + runner = testing_utils.InMemoryRunner( + app=App(name=request.function.__name__, root_agent=agent) + ) + + with pytest.raises(ValueError, match='kaboom 1'): + await runner.run_async(testing_utils.get_user_content('go')) + + +@pytest.mark.asyncio +async def test_node_tool_function_node_none_arg_for_non_optional_returns_error_dict( + request: pytest.FixtureRequest, +): + """Passing None to a non-optional parameter returns validation error dict.""" + + def calculate(x: int) -> int: + return x + + calc_node = FunctionNode(func=calculate) + agent = LlmAgent( + name='agent', + model=testing_utils.MockModel.create( + responses=[ + types.Part.from_function_call(name='calculate', args={'x': None}), + types.Part.from_text(text='done'), + ] + ), + tools=[calc_node], + ) + runner = testing_utils.InMemoryRunner( + app=App(name=request.function.__name__, root_agent=agent) + ) + + events = await runner.run_async(testing_utils.get_user_content('go')) + + responses = _function_responses(events) + assert len(responses) == 1 + assert set(responses[0]) == {'error'} + assert 'argument validation errors' in responses[0]['error'] + + +@pytest.mark.asyncio +async def test_node_tool_function_node_missing_node_input_param_returns_error_dict( + request: pytest.FixtureRequest, +): + """Missing a required node_input parameter returns validation error dict.""" + + def process(node_input: str) -> str: + return node_input + + proc_node = FunctionNode(func=process) + agent = LlmAgent( + name='agent', + model=testing_utils.MockModel.create( + responses=[ + types.Part.from_function_call(name='process', args={}), + types.Part.from_text(text='done'), + ] + ), + tools=[proc_node], + ) + runner = testing_utils.InMemoryRunner( + app=App(name=request.function.__name__, root_agent=agent) + ) + + events = await runner.run_async(testing_utils.get_user_content('go')) + + responses = _function_responses(events) + assert len(responses) == 1 + assert set(responses[0]) == {'error'} + assert 'argument validation errors' in responses[0]['error'] + assert 'Missing value for parameter "node_input"' in responses[0]['error'] + + +@pytest.mark.asyncio +async def test_node_tool_non_dict_input_schema_happy_path( + request: pytest.FixtureRequest, +): + """FunctionNode tool with non-dict input_schema successfully coerces node_input.""" + + def typed(node_input: _TypedInput) -> int: + return node_input.x * 2 + + typed_node = FunctionNode(func=typed) + typed_node.input_schema = _TypedInput + agent = LlmAgent( + name='agent', + model=testing_utils.MockModel.create( + responses=[ + types.Part.from_function_call(name='typed', args={'x': 21}), + types.Part.from_text(text='done'), + ] + ), + tools=[typed_node], + ) + runner = testing_utils.InMemoryRunner( + app=App(name=request.function.__name__, root_agent=agent) + ) + + events = await runner.run_async(testing_utils.get_user_content('go')) + + responses = _function_responses(events) + assert responses == [{'result': 42}] From e738c26fe5abcecffe2fbeba31f8d837086aad8a Mon Sep 17 00:00:00 2001 From: Kathy Wu Date: Tue, 29 Sep 2026 11:05:38 -0700 Subject: [PATCH 13/29] feat(mcp): add an opt-in modern-protocol connect path for MCP SDK 2.x ADK always brought a session up with `initialize()`, which pins the connection to the 2025 wire for its whole life even on SDK 2.x: the server issues an `Mcp-Session-Id` to route later requests back to one instance, and `clientInfo` is stated once rather than on each request. Setting `ADK_ENABLE_MCP_MODERN_PROTOCOL=1` probes `server/discover` first and falls back to the handshake on anything that is not a modern server, so a 2025-era server is unaffected. Off by default, and a no-op on SDK 1.x. Co-authored-by: Kathy Wu PiperOrigin-RevId: 990414661 --- src/google/adk/dependencies/_mcp.py | 20 +++++ src/google/adk/features/_feature_registry.py | 6 ++ .../adk/tools/mcp_tool/session_context.py | 37 ++++++++- .../tools/mcp_tool/test_dependencies_mcp.py | 15 +++- .../tools/mcp_tool/test_session_context.py | 76 +++++++++++++++++++ 5 files changed, 150 insertions(+), 4 deletions(-) diff --git a/src/google/adk/dependencies/_mcp.py b/src/google/adk/dependencies/_mcp.py index 6e49138016f..df0576fb38e 100644 --- a/src/google/adk/dependencies/_mcp.py +++ b/src/google/adk/dependencies/_mcp.py @@ -38,6 +38,9 @@ from __future__ import annotations +from typing import Awaitable +from typing import Callable + from mcp import ClientSession as ClientSession from mcp import SamplingCapability as SamplingCapability from mcp import StdioServerParameters as StdioServerParameters @@ -71,6 +74,22 @@ IS_MCP_SDK_V2 = False +# Tries `server/discover` and falls back to `initialize`. SDK 2.x only exposes +# this through its high-level `Client`, so we import the private function. +# Separate try so that if it moves, we don't misdetect the SDK as 1.x. +# `None` means connections use the legacy handshake. +if IS_MCP_SDK_V2: + try: + from mcp.client._probe import negotiate_auto as _negotiate_auto + except ImportError: + _negotiate_auto = None +else: + _negotiate_auto = None + +negotiate_auto: Callable[[ClientSession], Awaitable[None]] | None = ( + _negotiate_auto +) + __all__ = [ "IS_MCP_SDK_V2", "ClientSession", @@ -86,6 +105,7 @@ "StdioServerParameters", "Tool", "create_mcp_http_client", + "negotiate_auto", "sse_client", "stdio_client", "streamable_http_client", diff --git a/src/google/adk/features/_feature_registry.py b/src/google/adk/features/_feature_registry.py index 87c9440f35b..9e1889b540f 100644 --- a/src/google/adk/features/_feature_registry.py +++ b/src/google/adk/features/_feature_registry.py @@ -59,6 +59,9 @@ class FeatureName(str, Enum): # enum member by name. Keeping it private avoids a backward-compat # obligation for what is intended as a temporary, internal kill-switch. _MCP_GRACEFUL_ERROR_HANDLING = "MCP_GRACEFUL_ERROR_HANDLING" + # Off by default since it changes wire behavior. Enable with + # `ADK_ENABLE_MCP_MODERN_PROTOCOL=1`. No effect on MCP SDK 1.x. + _MCP_MODERN_PROTOCOL = "MCP_MODERN_PROTOCOL" MONGODB_TOOLSET = "MONGODB_TOOLSET" MONGODB_TOOL_SETTINGS = "MONGODB_TOOL_SETTINGS" PROGRESSIVE_SSE_STREAMING = "PROGRESSIVE_SSE_STREAMING" @@ -191,6 +194,9 @@ class FeatureConfig: FeatureName._MCP_GRACEFUL_ERROR_HANDLING: FeatureConfig( FeatureStage.EXPERIMENTAL, default_on=True ), + FeatureName._MCP_MODERN_PROTOCOL: FeatureConfig( + FeatureStage.EXPERIMENTAL, default_on=False + ), FeatureName.MONGODB_TOOLSET: FeatureConfig( FeatureStage.EXPERIMENTAL, default_on=True ), diff --git a/src/google/adk/tools/mcp_tool/session_context.py b/src/google/adk/tools/mcp_tool/session_context.py index fb5cf052b04..95d1bf09d77 100644 --- a/src/google/adk/tools/mcp_tool/session_context.py +++ b/src/google/adk/tools/mcp_tool/session_context.py @@ -18,6 +18,7 @@ from contextlib import AbstractAsyncContextManager from contextlib import AsyncExitStack from datetime import timedelta +import functools import logging from types import TracebackType from typing import Any @@ -28,6 +29,7 @@ from ...dependencies._mcp import ClientSession from ...dependencies._mcp import ElicitationFnT from ...dependencies._mcp import IS_MCP_SDK_V2 +from ...dependencies._mcp import negotiate_auto from ...dependencies._mcp import SamplingCapability from ...dependencies._mcp import SamplingFnT from ...dependencies._mcp import types @@ -45,6 +47,35 @@ _CLIENT_INFO = types.Implementation(name='google-adk', version=__version__) +async def _connect(session: ClientSession) -> None: + """Connects ``session`` using the newest protocol both sides support. + + When enabled, tries the modern protocol (``server/discover``) first, which + needs no ``Mcp-Session-Id`` and sends ``clientInfo`` on every request. Falls + back to ``initialize`` for older servers. + + Args: + session: The session to bring up. + """ + # pylint: disable-next=protected-access + if not is_feature_enabled(FeatureName._MCP_MODERN_PROTOCOL): + await session.initialize() + elif negotiate_auto is None: + _warn_probe_unavailable() + await session.initialize() + else: + await negotiate_auto(session) + + +# Cached so it only logs once. +@functools.cache +def _warn_probe_unavailable() -> None: + logger.warning( + 'MCP_MODERN_PROTOCOL is enabled, but the installed MCP SDK has no era' + ' probe ADK can use; MCP connections will use the legacy handshake.' + ) + + def _read_timeout(seconds: Optional[float]) -> Optional[float | timedelta]: """Converts a timeout in seconds to the type ``ClientSession`` expects. @@ -397,15 +428,15 @@ async def _run(self) -> None: ) # pylint: disable-next=protected-access if is_feature_enabled(FeatureName._MCP_GRACEFUL_ERROR_HANDLING): - # Use anyio.fail_after to keep session.initialize within the AnyIO + # Use anyio.fail_after to keep the handshake within the AnyIO # cancel scope instead of asyncio.wait_for which runs in a nested # task. import anyio with anyio.fail_after(self._timeout): - await session.initialize() + await _connect(session) else: - await asyncio.wait_for(session.initialize(), timeout=self._timeout) + await asyncio.wait_for(_connect(session), timeout=self._timeout) logger.debug('Session has been successfully initialized') self._session = session diff --git a/tests/unittests/tools/mcp_tool/test_dependencies_mcp.py b/tests/unittests/tools/mcp_tool/test_dependencies_mcp.py index cf10a1ea9f5..3ad9e740791 100644 --- a/tests/unittests/tools/mcp_tool/test_dependencies_mcp.py +++ b/tests/unittests/tools/mcp_tool/test_dependencies_mcp.py @@ -76,6 +76,10 @@ def _sdk_imports(source: str) -> list[str]: return lines +# Exports that are `None` on SDK 1.x because they only exist in 2.x. +_OPTIONAL_NAMES = frozenset({'negotiate_auto'}) + + class TestTheSeamHolds: """The seam is only worth having if nothing routes around it.""" @@ -102,11 +106,20 @@ def test_every_advertised_name_resolves(self): missing = [ name for name in mcp_dependency.__all__ - if getattr(mcp_dependency, name, None) is None + if name not in _OPTIONAL_NAMES + and getattr(mcp_dependency, name, None) is None ] assert not missing + def test_every_optional_name_is_bound(self): + """Optional names can be `None`, but must still be defined.""" + unbound = [ + name for name in _OPTIONAL_NAMES if not hasattr(mcp_dependency, name) + ] + + assert not unbound + def test_the_advertised_sdk_name_is_the_one_the_seam_imports(self): """Telemetry looks the SDK up in `sys.modules` by this name. diff --git a/tests/unittests/tools/mcp_tool/test_session_context.py b/tests/unittests/tools/mcp_tool/test_session_context.py index 899885800de..2f4dda6cba0 100644 --- a/tests/unittests/tools/mcp_tool/test_session_context.py +++ b/tests/unittests/tools/mcp_tool/test_session_context.py @@ -17,6 +17,7 @@ import asyncio from contextlib import AsyncExitStack from datetime import timedelta +import logging import time from unittest.mock import AsyncMock from unittest.mock import Mock @@ -24,8 +25,10 @@ from google.adk.features import FeatureName from google.adk.features._feature_registry import temporary_feature_override +from google.adk.tools.mcp_tool.session_context import _connect from google.adk.tools.mcp_tool.session_context import _format_exception from google.adk.tools.mcp_tool.session_context import _read_timeout +from google.adk.tools.mcp_tool.session_context import _warn_probe_unavailable from google.adk.tools.mcp_tool.session_context import SessionContext from google.adk.version import __version__ import httpx @@ -748,6 +751,79 @@ async def test_names_adk_in_client_info(self, is_stdio): assert kwargs['client_info'].version == __version__ +class TestConnect: + """Tests for `_connect`.""" + + @pytest.fixture(autouse=True) + def _reset_probe_warning(self): + _warn_probe_unavailable.cache_clear() + + @pytest.mark.asyncio + async def test_uses_handshake_by_default(self): + """Uses `initialize` when the flag is off.""" + session = Mock() + session.initialize = AsyncMock() + probe = AsyncMock() + + with patch( + 'google.adk.tools.mcp_tool.session_context.negotiate_auto', probe + ): + await _connect(session) + + session.initialize.assert_awaited_once() + probe.assert_not_awaited() + + @pytest.mark.asyncio + async def test_probes_when_enabled(self): + """Uses `negotiate_auto` when the flag is on.""" + session = Mock() + session.initialize = AsyncMock() + probe = AsyncMock() + + with ( + patch( + 'google.adk.tools.mcp_tool.session_context.negotiate_auto', probe + ), + temporary_feature_override(FeatureName._MCP_MODERN_PROTOCOL, True), + ): + await _connect(session) + + probe.assert_awaited_once_with(session) + session.initialize.assert_not_awaited() + + @pytest.mark.asyncio + async def test_falls_back_when_probe_unavailable(self, caplog): + """Falls back to `initialize` and warns once if the SDK lacks the probe.""" + session = Mock() + session.initialize = AsyncMock() + + with ( + patch('google.adk.tools.mcp_tool.session_context.negotiate_auto', None), + temporary_feature_override(FeatureName._MCP_MODERN_PROTOCOL, True), + caplog.at_level(logging.WARNING), + ): + await _connect(session) + await _connect(session) + + assert session.initialize.await_count == 2 + assert caplog.text.count('no era probe') == 1 + + @pytest.mark.asyncio + async def test_no_probe_warning_when_disabled(self, caplog): + """No warning about a missing probe when the flag is off.""" + session = Mock() + session.initialize = AsyncMock() + + with ( + patch('google.adk.tools.mcp_tool.session_context.negotiate_auto', None), + caplog.at_level(logging.WARNING), + ): + await _connect(session) + + session.initialize.assert_awaited_once() + assert 'no era probe' not in caplog.text + + class TestSessionContextIsTaskAlive: """Tests for the SessionContext._is_task_alive property.""" From 5a0421cbc58ae7b9e00ae4cd41f124863b0cbb70 Mon Sep 17 00:00:00 2001 From: Shangjie Chen Date: Tue, 29 Sep 2026 11:38:37 -0700 Subject: [PATCH 14/29] feat(workflow): propagate skip_summarization from node tools to tool response Propagate `ctx.actions.skip_summarization` from dynamically executed child nodes to the parent context in `run_node_internal`, and include `NodeTool` alongside `AgentTool` when attaching displayable tool output to `function_response` events with `skip_summarization=True`. This allows `@node` and `Workflow` tools to set `ctx.actions.skip_summarization = True` inside the node so their terminal output is emitted directly as the final user-visible text response without triggering a follow-up LLM summarization turn, while keeping `NodeTool` internal. Co-authored-by: Shangjie Chen PiperOrigin-RevId: 990434765 --- .../adk/flows/llm_flows/tools/_caller.py | 16 +-- .../adk/workflow/_dynamic_node_scheduler.py | 2 + tests/unittests/workflow/test_node_tool.py | 99 +++++++++++++++++++ 3 files changed, 110 insertions(+), 7 deletions(-) diff --git a/src/google/adk/flows/llm_flows/tools/_caller.py b/src/google/adk/flows/llm_flows/tools/_caller.py index c99e1a07bee..69c754fc8c2 100644 --- a/src/google/adk/flows/llm_flows/tools/_caller.py +++ b/src/google/adk/flows/llm_flows/tools/_caller.py @@ -542,17 +542,19 @@ def _build_response_event( and 'error' not in function_result and has_displayable_result ): - # Imported lazily: AgentTool is only needed on the skip-summarization - # path, so it is not worth pulling into every functions.py import. + # Imported lazily: AgentTool and NodeTool are only needed on the + # skip-summarization path, so they are not worth pulling into every + # functions.py import. + from ....tools._node_tool import NodeTool from ....tools.agent_tool import AgentTool - # This is scoped to AgentTool deliberately: other tools (e.g. UI/widget- - # rendering tools) set skip_summarization precisely because their function - # response is an internal acknowledgement that must NOT be surfaced as - # visible text. AgentTool subclasses can still return None (e.g. + # This is scoped to AgentTool and NodeTool deliberately: other tools (e.g. + # UI/widget-rendering tools) set skip_summarization precisely because their + # function response is an internal acknowledgement that must NOT be surfaced + # as visible text. AgentTool subclasses can still return None (e.g. # _SingleTurnAgentTool delegating to run_node), hence the # has_displayable_result guard above. - if isinstance(tool, AgentTool): + if isinstance(tool, (AgentTool, NodeTool)): if isinstance(display_result, str): result_text = display_result else: diff --git a/src/google/adk/workflow/_dynamic_node_scheduler.py b/src/google/adk/workflow/_dynamic_node_scheduler.py index 304db1b955e..25a868be3b0 100644 --- a/src/google/adk/workflow/_dynamic_node_scheduler.py +++ b/src/google/adk/workflow/_dynamic_node_scheduler.py @@ -731,6 +731,8 @@ async def run_node_internal( ) transfer_to_agent = child_ctx.actions.transfer_to_agent if child_ctx else None + if child_ctx and child_ctx.actions.skip_summarization: + ctx.actions.skip_summarization = True if not return_ctx: if child_ctx.error: diff --git a/tests/unittests/workflow/test_node_tool.py b/tests/unittests/workflow/test_node_tool.py index f16a3da19ea..a7371ec4b40 100644 --- a/tests/unittests/workflow/test_node_tool.py +++ b/tests/unittests/workflow/test_node_tool.py @@ -1826,3 +1826,102 @@ def typed(node_input: _TypedInput) -> int: responses = _function_responses(events) assert responses == [{'result': 42}] + + +@pytest.mark.asyncio +async def test_node_tool_skip_summarization_returns_workflow_output_directly( + request: pytest.FixtureRequest, +): + """Setting ctx.actions.skip_summarization=True inside a workflow node tool emits output directly as final text without a second LLM turn.""" + + def run_parallel_report(node_input: GreetRequest, ctx: Context) -> str: + ctx.actions.skip_summarization = True + return f'Exact workflow report for {node_input.request}' + + sub_workflow = Workflow( + name='report_workflow', + description='Generates an exact report.', + input_schema=GreetRequest, + edges=[(START, run_parallel_report)], + ) + + mock_model = testing_utils.MockModel.create( + responses=[ + types.Part.from_function_call( + name='report_workflow', + args={'request': 'Project X'}, + ), + types.Part.from_text( + text='Should never be reached because summarization is skipped.' + ), + ] + ) + + parent_agent = LlmAgent( + name='parent_agent', + model=mock_model, + tools=[sub_workflow], + ) + + app = App(name=request.function.__name__, root_agent=parent_agent) + runner = testing_utils.InMemoryRunner(app=app) + + events = await runner.run_async(testing_utils.get_user_content('Report X')) + + parent_final_events = [ + e for e in events if e.author == 'parent_agent' and e.is_final_response() + ] + assert len(parent_final_events) == 1 + last_event = events[-1] + assert last_event == parent_final_events[0] + assert last_event.actions.skip_summarization is True + assert any(p.function_response for p in last_event.content.parts) + text_parts = [p.text for p in last_event.content.parts if p.text] + assert text_parts == ['Exact workflow report for Project X'] + assert len(mock_model.requests) == 1 + + +@pytest.mark.asyncio +async def test_function_node_tool_skip_summarization_returns_output_directly( + request: pytest.FixtureRequest, +): + """Setting ctx.actions.skip_summarization=True inside a @node tool emits output directly as final text.""" + + @node + def generate_report(project: str, ctx: Context) -> str: + """Generates a report for a project.""" + ctx.actions.skip_summarization = True + return f'Report for {project}: Ready' + + mock_model = testing_utils.MockModel.create( + responses=[ + types.Part.from_function_call( + name='generate_report', + args={'project': 'Apollo'}, + ), + types.Part.from_text( + text='Should never be reached because summarization is skipped.' + ), + ] + ) + + parent_agent = LlmAgent( + name='parent_agent', + model=mock_model, + tools=[generate_report], + ) + + app = App(name=request.function.__name__, root_agent=parent_agent) + runner = testing_utils.InMemoryRunner(app=app) + + events = await runner.run_async( + testing_utils.get_user_content('Report Apollo') + ) + + last_event = events[-1] + assert last_event.author == 'parent_agent' + assert last_event.is_final_response() + assert last_event.actions.skip_summarization is True + text_parts = [p.text for p in last_event.content.parts if p.text] + assert text_parts == ['Report for Apollo: Ready'] + assert len(mock_model.requests) == 1 From fd2ca8773ed51ec13a96670fdeb11f333727217a Mon Sep 17 00:00:00 2001 From: George Weale Date: Tue, 29 Sep 2026 12:25:13 -0700 Subject: [PATCH 15/29] fix(eval): skip content-less events when mapping Vertex multi-turn turns Invocations now keep the final event with its content removed so the efficiency metrics can read its token usage, and the Vertex multi-turn facade sent that event as an empty agent message in every turn. Skips intermediate events without content when mapping a turn. Co-authored-by: George Weale PiperOrigin-RevId: 990461145 --- .../adk/evaluation/vertex_ai_eval_facade.py | 2 ++ .../evaluation/test_vertex_ai_eval_facade.py | 16 ++++++++++++++++ 2 files changed, 18 insertions(+) diff --git a/src/google/adk/evaluation/vertex_ai_eval_facade.py b/src/google/adk/evaluation/vertex_ai_eval_facade.py index 69d2b82f329..124df2c04cf 100644 --- a/src/google/adk/evaluation/vertex_ai_eval_facade.py +++ b/src/google/adk/evaluation/vertex_ai_eval_facade.py @@ -325,6 +325,8 @@ def _map_invocation_turn( if isinstance(invocation.intermediate_data, InvocationEvents): for invocation_event in invocation.intermediate_data.invocation_events: + if invocation_event.content is None: + continue agent_events.append( _MultiTurnVertexiAiEvalFacade._map_inovcation_event_to_agent_event( invocation_event diff --git a/tests/unittests/evaluation/test_vertex_ai_eval_facade.py b/tests/unittests/evaluation/test_vertex_ai_eval_facade.py index 285f8c70629..4720e825d5d 100644 --- a/tests/unittests/evaluation/test_vertex_ai_eval_facade.py +++ b/tests/unittests/evaluation/test_vertex_ai_eval_facade.py @@ -454,6 +454,22 @@ def test_map_invocation_turn(self): assert conversation_turn.events[2].author == "agent" assert conversation_turn.events[2].content.parts[0].text == "final response" + def test_map_invocation_turn_skips_events_without_content(self): + invocation = Invocation( + invocation_id="inv1", + user_content=genai_types.Content(parts=[genai_types.Part(text="hi")]), + intermediate_data=InvocationEvents( + invocation_events=[InvocationEvent(author="agent1", content=None)] + ), + final_response=genai_types.Content( + parts=[genai_types.Part(text="hello")] + ), + ) + conversation_turn = _MultiTurnVertexiAiEvalFacade._map_invocation_turn( + 0, invocation + ) + assert [e.author for e in conversation_turn.events] == ["user", "agent"] + def test_get_turns(self): invocations = [ Invocation( From 86329806d1eb5c4cc767ab291730ccdf338d5ae5 Mon Sep 17 00:00:00 2001 From: George Weale Date: Tue, 29 Sep 2026 12:47:14 -0700 Subject: [PATCH 16/29] fix: detect a dead MCP session whose transport sits behind a dispatcher MCP SDK 2.x moved the transport off the client session and behind a dispatcher, so ADK's liveness check read attributes that no longer exist and a pooled session whose server had died looked healthy forever, failing every later call. A session with no streams of its own is now checked through its dispatcher's closed flag, which restores the existing reconnect path on both SDK majors. Co-authored-by: George Weale PiperOrigin-RevId: 990472674 --- .../adk/tools/mcp_tool/mcp_session_manager.py | 29 +++++-- .../mcp_tool/test_mcp_session_manager.py | 85 ++++++++++++++++++- 2 files changed, 104 insertions(+), 10 deletions(-) diff --git a/src/google/adk/tools/mcp_tool/mcp_session_manager.py b/src/google/adk/tools/mcp_tool/mcp_session_manager.py index 84d262c0981..c83fbcda12c 100644 --- a/src/google/adk/tools/mcp_tool/mcp_session_manager.py +++ b/src/google/adk/tools/mcp_tool/mcp_session_manager.py @@ -1006,15 +1006,20 @@ def _merge_headers( def _is_session_disconnected(self, session: ClientSession) -> bool: """Checks if a session is disconnected or closed. - Reads two attributes ADK does not own: the SDK holds the transport streams - on the session privately, and each stream reports its own closed flag. A - session that lacks either one reads as connected rather than raising, - because a release is free to restructure both away and this probe is not - the only thing standing between a dead session and a caller. + Reads attributes ADK does not own, and where they hang moved between SDK + majors. On 1.x the session holds the transport streams and each stream + reports its own closed flag. On 2.x the transport moved behind a + dispatcher, which reports one closed flag of its own and need not hold + streams at all, so a session holding no streams is read there instead. A + session offering neither reads as connected rather than raising, because + a release is free to restructure them away and this probe is not the only + thing standing between a dead session and a caller. `create_session` pairs this with `SessionContext._is_task_alive`, which - ADK owns and which catches strictly more: a crashed transport can leave - the streams open while the task behind them is already dead. That pairing + ADK owns. Neither check subsumes the other: a crashed transport can + leave the streams open while the task behind them is already dead, and a + transport that closes under a live session leaves that task parked on + its close event, where only these flags report the death. That pairing runs under `_MCP_GRACEFUL_ERROR_HANDLING`, which is on by default. The kill switch drops it and leaves this probe on its own. @@ -1027,6 +1032,16 @@ def _is_session_disconnected(self, session: ClientSession) -> bool: Returns: True if the session is known to be disconnected, False otherwise. """ + if not hasattr(session, '_read_stream'): + dispatcher = getattr(session, '_dispatcher', None) + if not hasattr(dispatcher, '_closed'): + logger.debug( + 'MCP session %s offers no closed flag to read, on itself or on a' + ' dispatcher; reading it as connected.', + type(session).__name__, + ) + return False + return bool(getattr(dispatcher, '_closed', False)) read_stream = getattr(session, '_read_stream', None) write_stream = getattr(session, '_write_stream', None) return bool( diff --git a/tests/unittests/tools/mcp_tool/test_mcp_session_manager.py b/tests/unittests/tools/mcp_tool/test_mcp_session_manager.py index ff02edbf983..9a172f8766a 100644 --- a/tests/unittests/tools/mcp_tool/test_mcp_session_manager.py +++ b/tests/unittests/tools/mcp_tool/test_mcp_session_manager.py @@ -25,7 +25,9 @@ from unittest.mock import patch import urllib.parse +import anyio from google.adk.dependencies import _httpx as httpx +from google.adk.dependencies._mcp import ClientSession from google.adk.dependencies._mcp import IS_MCP_SDK_V2 from google.adk.dependencies._mcp import McpError from google.adk.features import FeatureName @@ -466,13 +468,14 @@ def test_is_session_disconnected_write_stream_closed(self): session._write_stream._closed = True assert manager._is_session_disconnected(session) - def test_is_session_disconnected_without_streams(self): + def test_is_session_disconnected_without_streams(self, caplog): """A session that holds no streams reads as connected, and does not raise. Both attributes are private to the SDK. A release is free to move the streams off `ClientSession`, and this must degrade to the `SessionContext` task check rather than take down every tool call with an - `AttributeError`. + `AttributeError`. It logs on the way, so the next SDK bump leaving this + probe nothing to read shows up instead of going quiet. The stand-in is a bare class on purpose: a `Mock` would answer to `_read_stream` and pass this vacuously. @@ -482,7 +485,12 @@ class SessionWithoutStreams: pass manager = MCPSessionManager(self.mock_stdio_connection_params) - assert not manager._is_session_disconnected(SessionWithoutStreams()) + with caplog.at_level(logging.DEBUG): + assert not manager._is_session_disconnected(SessionWithoutStreams()) + assert any( + "SessionWithoutStreams" in record.getMessage() + for record in caplog.records + ) def test_is_session_disconnected_with_streams_that_have_no_flag(self): """A stream that stops reporting a closed flag reads as connected too.""" @@ -499,6 +507,77 @@ def __init__(self): manager = MCPSessionManager(self.mock_stdio_connection_params) assert not manager._is_session_disconnected(SessionWithBareStreams()) + def test_is_session_disconnected_reads_a_dispatcher_closed_flag(self): + """A session whose transport sits behind a dispatcher is still probed. + + The SDK moved the transport off `ClientSession` and behind a dispatcher in + its 2.x line. Looking only at the session reads a dead transport as live + and leaves the pooled session wedged for every later call. The dispatcher + need not hold streams at all, so its own flag is what gets read. + """ + + class Dispatcher: + + def __init__(self): + self._closed = False + + class SessionWithDispatcher: + + def __init__(self): + self._dispatcher = Dispatcher() + + manager = MCPSessionManager(self.mock_stdio_connection_params) + + session = SessionWithDispatcher() + assert not manager._is_session_disconnected(session) + + session._dispatcher._closed = True + assert manager._is_session_disconnected(session) + + def test_is_session_disconnected_prefers_the_session_over_a_dispatcher(self): + """A session holding its own streams is read there, dispatcher or not.""" + + class Stream: + + def __init__(self): + self._closed = False + + class Dispatcher: + + def __init__(self): + self._closed = True + + class SessionWithBoth: + + def __init__(self): + self._read_stream = Stream() + self._write_stream = Stream() + self._dispatcher = Dispatcher() + + manager = MCPSessionManager(self.mock_stdio_connection_params) + assert not manager._is_session_disconnected(SessionWithBoth()) + + def test_is_session_disconnected_reads_a_real_client_session(self): + """The probe finds its flag on a real session, not only on a stand-in. + + The classes above are written here, so they prove the branching and not + the layout. This one builds the installed SDK's own `ClientSession` and + fails if the attribute the probe reads is not where it looks. + """ + write_stream, read_stream = anyio.create_memory_object_stream(1) + session = ClientSession(read_stream, write_stream) + + manager = MCPSessionManager(self.mock_stdio_connection_params) + assert not manager._is_session_disconnected(session) + + if hasattr(session, "_read_stream"): + assert hasattr(session._read_stream, "_closed") + session._read_stream._closed = True + else: + assert hasattr(session._dispatcher, "_closed") + session._dispatcher._closed = True + assert manager._is_session_disconnected(session) + @pytest.mark.asyncio async def test_discard_session_drops_a_session_that_still_looks_healthy(self): """The pooled session goes even though its streams report open.""" From 4d241bff1bdf63ddf5ef947a2aa5d145b008c013 Mon Sep 17 00:00:00 2001 From: Liang Wu Date: Tue, 29 Sep 2026 13:33:27 -0700 Subject: [PATCH 17/29] fix: only apply --avatar_config to live sessions requesting video The server-wide avatar configuration added to `adk web` and `adk api_server` was attached to every /run_live session, including the default audio-only ones. Avatars are rendered as video, so only set `RunConfig.avatar_config` when the client requests the VIDEO modality, and say so in the `--avatar_config` help text. Also adds a test for the unreadable avatar configuration file path. Co-authored-by: Liang Wu PiperOrigin-RevId: 990498163 --- src/google/adk/cli/api_server.py | 6 +- src/google/adk/cli/cli_tools_click.py | 4 +- src/google/adk/cli/fast_api.py | 3 +- .../cli/test_adk_web_server_run_live.py | 72 ++++++++++++++++++- .../cli/utils/test_cli_tools_click.py | 14 ++++ 5 files changed, 94 insertions(+), 5 deletions(-) diff --git a/src/google/adk/cli/api_server.py b/src/google/adk/cli/api_server.py index 17739c70212..8a50bc8ee76 100644 --- a/src/google/adk/cli/api_server.py +++ b/src/google/adk/cli/api_server.py @@ -2171,7 +2171,11 @@ async def forward_events(): ), save_live_blob=save_live_blob, explicit_vad_signal=explicit_vad_signal, - avatar_config=self.avatar_config, + # Avatars are rendered as video, so only apply the server-wide + # avatar config to sessions that request VIDEO output. + avatar_config=( + self.avatar_config if "VIDEO" in modalities else None + ), ) async with Aclosing( runner.run_live( diff --git a/src/google/adk/cli/cli_tools_click.py b/src/google/adk/cli/cli_tools_click.py index 8bf1d15c60d..3b22f763df7 100644 --- a/src/google/adk/cli/cli_tools_click.py +++ b/src/google/adk/cli/cli_tools_click.py @@ -2086,7 +2086,9 @@ def decorator(func): callback=_parse_avatar_config, help=( "Optional. AvatarConfig as an inline JSON object or a path to a" - " JSON file. Applied to live sessions." + " JSON file. Applied only to /run_live sessions whose client" + " requests video output (modalities=VIDEO); other live sessions" + " ignore it." ), default=None, ) diff --git a/src/google/adk/cli/fast_api.py b/src/google/adk/cli/fast_api.py index fde0d6c13e3..6725e59392c 100644 --- a/src/google/adk/cli/fast_api.py +++ b/src/google/adk/cli/fast_api.py @@ -195,7 +195,8 @@ def get_fast_api_app( gemini_enterprise_app_name: The Gemini Enterprise app name to use for the agent. express_mode: Whether to enable express mode. - avatar_config: Avatar configuration to apply to live agent runs. + avatar_config: Avatar configuration to apply to live agent runs that + request VIDEO output. Returns: The configured FastAPI application instance. diff --git a/tests/unittests/cli/test_adk_web_server_run_live.py b/tests/unittests/cli/test_adk_web_server_run_live.py index f5cd6c66ca1..f51e6c1bc6d 100644 --- a/tests/unittests/cli/test_adk_web_server_run_live.py +++ b/tests/unittests/cli/test_adk_web_server_run_live.py @@ -124,8 +124,76 @@ async def _get_runner_async(_self, _app_name: str): assert run_config.session_resumption.transparent is True assert run_config.save_live_blob is True assert run_config.explicit_vad_signal is True - assert run_config.avatar_config is not None - assert run_config.avatar_config.avatar_name == "Kai" + # No VIDEO modality was requested, so the server avatar config is skipped. + assert run_config.avatar_config is None + + +@pytest.mark.parametrize( + ("modalities_query", "expect_avatar"), + [ + ("&modalities=VIDEO", True), + ("&modalities=AUDIO&modalities=VIDEO", True), + ("&modalities=AUDIO", False), + ("&modalities=TEXT", False), + ("", False), + ], +) +def test_run_live_applies_avatar_config_only_for_video( + modalities_query: str, expect_avatar: bool +): + """The server avatar config is sent only when VIDEO output is requested.""" + session_service = InMemorySessionService() + asyncio.run( + session_service.create_session( + app_name="test_app", + user_id="user", + session_id="session", + state={}, + ) + ) + + runner = _CapturingRunner() + adk_web_server = AdkWebServer( + agent_loader=_DummyAgentLoader(), + session_service=session_service, + memory_service=types.SimpleNamespace(), + artifact_service=types.SimpleNamespace(), + credential_service=types.SimpleNamespace(), + eval_sets_manager=types.SimpleNamespace(), + eval_set_results_manager=types.SimpleNamespace(), + agents_dir=".", + avatar_config=genai_types.AvatarConfig(avatar_name="Kai"), + ) + + async def _get_runner_async(_self, _app_name: str): + return runner + + adk_web_server.get_runner_async = _get_runner_async.__get__(adk_web_server) # pytype: disable=attribute-error + + fast_api_app = adk_web_server.get_fast_api_app( + setup_observer=lambda _observer, _server: None, + tear_down_observer=lambda _observer, _server: None, + ) + + client = TestClient(fast_api_app) + url = ( + "/run_live" + "?app_name=test_app" + "&user_id=user" + "&session_id=session" + f"{modalities_query}" + ) + + with client.websocket_connect(url) as ws: + _ = ws.receive_text() + + run_config = runner.captured_run_config + assert run_config is not None + if expect_avatar: + assert run_config.avatar_config is not None + assert run_config.avatar_config.avatar_name == "Kai" + else: + assert run_config.avatar_config is None @pytest.mark.parametrize( diff --git a/tests/unittests/cli/utils/test_cli_tools_click.py b/tests/unittests/cli/utils/test_cli_tools_click.py index 093a2ff49c4..b527ce0648c 100644 --- a/tests/unittests/cli/utils/test_cli_tools_click.py +++ b/tests/unittests/cli/utils/test_cli_tools_click.py @@ -2835,6 +2835,20 @@ def test_fast_api_common_options_rejects_invalid_avatar_config() -> None: assert "valid AvatarConfig JSON object" in result.output +def test_fast_api_common_options_rejects_missing_avatar_config_file( + tmp_path: Path, +) -> None: + """A non-JSON value that is not a readable file is a usage error.""" + command, _ = _fast_api_command() + missing_path = tmp_path / "missing_avatar.json" + + result = CliRunner().invoke(command, ["--avatar_config", str(missing_path)]) + + assert result.exit_code == 2 + assert "could not read avatar configuration file" in result.output + assert "missing_avatar.json" in result.output + + # adk test @pytest.fixture def fake_pytest_run(monkeypatch: pytest.MonkeyPatch): From fd14aec26534adb9743e963f39af468f1f0c329f Mon Sep 17 00:00:00 2001 From: Shinobu Aoki Date: Tue, 29 Sep 2026 13:49:31 -0700 Subject: [PATCH 18/29] fix: follow redirects when downloading skills in GcpSkillRegistry Merge https://github.com/google/adk-python/pull/6824 PiperOrigin-RevId: 990507320 --- .../skill_registry/gcp_skill_registry.py | 26 ++++- .../skill_registry/test_gcp_skill_registry.py | 103 +++++++++++++++++- 2 files changed, 126 insertions(+), 3 deletions(-) diff --git a/src/google/adk/integrations/skill_registry/gcp_skill_registry.py b/src/google/adk/integrations/skill_registry/gcp_skill_registry.py index e25a55a848e..2eaf14db292 100644 --- a/src/google/adk/integrations/skill_registry/gcp_skill_registry.py +++ b/src/google/adk/integrations/skill_registry/gcp_skill_registry.py @@ -184,9 +184,31 @@ async def _make_request( def _create_httpx_client(self) -> httpx.AsyncClient: """Creates a new httpx.AsyncClient with appropriate SSL/mTLS configuration.""" + base_host = httpx.URL(self.base_url).host + + async def _drop_cross_origin_goog_headers(request: httpx.Request) -> None: + if request.url.host != base_host: + for header in list(request.headers): + if header.lower().startswith("x-goog-"): + del request.headers[header] + + # The Agent Registry media download (alt=media) replies with a 302 to a + # short-lived GCS signed URL, so the client must follow redirects; httpx + # drops the Authorization header on cross-origin redirects, but retains + # custom headers like x-goog-user-project and x-goog-api-client. GCS + # requires all x-goog-* headers on a signed request to match its signature, + # so we drop them when redirected off the base API host. + event_hooks = {"request": [_drop_cross_origin_goog_headers]} if self._ssl_context is not None: - return httpx.AsyncClient(verify=self._ssl_context) - return httpx.AsyncClient() + return httpx.AsyncClient( + verify=self._ssl_context, + follow_redirects=True, + event_hooks=event_hooks, + ) + return httpx.AsyncClient( + follow_redirects=True, + event_hooks=event_hooks, + ) async def get_skill(self, *, name: str) -> models.Skill: """Fetches a skill from the registry. diff --git a/tests/unittests/integrations/skill_registry/test_gcp_skill_registry.py b/tests/unittests/integrations/skill_registry/test_gcp_skill_registry.py index 61643944786..aaef0e0df78 100644 --- a/tests/unittests/integrations/skill_registry/test_gcp_skill_registry.py +++ b/tests/unittests/integrations/skill_registry/test_gcp_skill_registry.py @@ -17,11 +17,13 @@ import io import logging import os +import ssl from unittest import mock import zipfile from google.adk.integrations.skill_registry import gcp_skill_registry from google.adk.utils._google_client_headers import merge_tracking_headers +import httpx import pytest @@ -588,7 +590,11 @@ async def mock_get(url, *unused_args, **kwargs): skill = await registry.get_skill(name="my-skill") # Verify AsyncClient was instantiated with verify=mock_ssl_context - mock_client_class.assert_called_with(verify=mock_ssl_context) + mock_client_class.assert_called_with( + verify=mock_ssl_context, + follow_redirects=True, + event_hooks=mock.ANY, + ) assert skill.frontmatter.name == "my-skill" @@ -652,3 +658,98 @@ async def test_search_skills_result_passes_frontmatter_validation(): results[0].model_dump() ) assert validated.name == "cloud.google.com-agent-platform-eval-flywheel" + + +@pytest.mark.asyncio +async def test_create_httpx_client_follows_redirects(): + """Clients follow the 302 redirect issued by the media download endpoint.""" + registry = gcp_skill_registry.GCPSkillRegistry() + + client = registry._create_httpx_client() + try: + assert client.follow_redirects is True + assert client.event_hooks["request"] + finally: + await client.aclose() + + registry._ssl_context = ssl.create_default_context() + client = registry._create_httpx_client() + try: + assert client.follow_redirects is True + assert client.event_hooks["request"] + finally: + await client.aclose() + + +@pytest.mark.asyncio +async def test_get_skill_drops_goog_headers_on_redirect(): + """Verifies that x-goog-* and auth headers are stripped on cross-origin redirects.""" + fake_zip = _create_fake_zip_bytes() + + def transport_handler(request: httpx.Request) -> httpx.Response: + if "skills/my-skill" in str(request.url) and "revisions" not in str( + request.url + ): + return httpx.Response( + 200, + json={ + "name": ( + "projects/test-project/locations/us-central1/skills/my-skill" + ), + "defaultRevision": ( + "projects/test-project/locations/us-central1/skills/my-skill/revisions/rev-123" + ), + }, + ) + if "alt=media" in str(request.url) or ( + request.url.params and request.url.params.get("alt") == "media" + ): + if request.url.host == "agentregistry.googleapis.com": + return httpx.Response( + 302, + headers={ + "Location": ( + "https://storage.googleapis.com/download/storage/v1/b/bucket/o/skill.zip?signature=123" + ) + }, + ) + if request.url.host == "storage.googleapis.com": + for header in request.headers: + if header.lower().startswith("x-goog-"): + return httpx.Response( + 403, + text=f"SignatureDoesNotMatch: Header {header} not signed", + ) + if header.lower() == "authorization": + return httpx.Response( + 403, + text="SignatureDoesNotMatch: Authorization not signed", + ) + return httpx.Response(200, content=fake_zip) + return httpx.Response(404, text=f"Not found: {request.url}") + + mock_creds = mock.MagicMock() + mock_creds.valid = True + mock_creds.token = "test-token" + mock_creds.quota_project_id = "test-quota" + + registry = gcp_skill_registry.GCPSkillRegistry( + project_id="test-project", + location="us-central1", + credentials=mock_creds, + ) + + orig_create = registry._create_httpx_client + + def custom_create(): + client = orig_create() + return httpx.AsyncClient( + transport=httpx.MockTransport(transport_handler), + follow_redirects=client.follow_redirects, + event_hooks=client.event_hooks, + ) + + registry._create_httpx_client = custom_create + + skill = await registry.get_skill(name="my-skill") + assert skill.frontmatter.name == "my-skill" From f7146ac3ce08c579f705193011bef91b919f0824 Mon Sep 17 00:00:00 2001 From: YASHcode-IIITV Date: Tue, 29 Sep 2026 14:37:47 -0700 Subject: [PATCH 19/29] test: cover api server auto create session Merge https://github.com/google/adk-python/pull/6786 PiperOrigin-RevId: 990538126 --- .../cli/utils/test_cli_tools_click.py | 29 +++++++++++++++++++ 1 file changed, 29 insertions(+) diff --git a/tests/unittests/cli/utils/test_cli_tools_click.py b/tests/unittests/cli/utils/test_cli_tools_click.py index b527ce0648c..97057e9faa6 100644 --- a/tests/unittests/cli/utils/test_cli_tools_click.py +++ b/tests/unittests/cli/utils/test_cli_tools_click.py @@ -1758,6 +1758,35 @@ def test_cli_web_passes_service_uris( assert called_kwargs.get("memory_service_uri") == "rag://mycorpus" +def test_cli_api_server_passes_auto_create_session( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, + _patch_uvicorn: _Recorder, +) -> None: + """`adk api_server --auto_create_session` enables automatic sessions.""" + agents_dir = tmp_path / "agents_api" + agents_dir.mkdir() + + mock_get_app = _Recorder() + monkeypatch.setattr("google.adk.cli.fast_api.get_fast_api_app", mock_get_app) + + runner = CliRunner() + result = runner.invoke( + cli_tools_click.main, + [ + "api_server", + str(agents_dir), + "--auto_create_session", + ], + ) + + assert result.exit_code == 0 + assert mock_get_app.calls + + called_kwargs = mock_get_app.calls[-1][1] + assert called_kwargs["auto_create_session"] is True + + @pytest.mark.parametrize("command", ["web", "api_server"]) @pytest.mark.parametrize("host", ["127.0.0.1", "0.0.0.0"]) def test_cli_arms_rebinding_guard_with_the_address_it_binds( From 09601045f4daced76b00f92c641fb1899ae3fa53 Mon Sep 17 00:00:00 2001 From: Chaitanya Laxman Date: Tue, 29 Sep 2026 14:45:49 -0700 Subject: [PATCH 20/29] feat: propagate grounding metadata from MCP _meta Merge https://github.com/google/adk-python/pull/7047 Fixes #6081 PiperOrigin-RevId: 990543079 --- .../adk/flows/llm_flows/core/_finalizer.py | 9 ++- src/google/adk/tools/mcp_tool/mcp_tool.py | 32 +++++++++ src/google/adk/tools/mcp_tool/mcp_toolset.py | 6 ++ .../flows/llm_flows/core/test_finalizer.py | 17 +++-- .../unittests/tools/mcp_tool/test_mcp_tool.py | 69 +++++++++++++++++++ .../tools/mcp_tool/test_mcp_toolset.py | 22 ++++++ 6 files changed, 145 insertions(+), 10 deletions(-) diff --git a/src/google/adk/flows/llm_flows/core/_finalizer.py b/src/google/adk/flows/llm_flows/core/_finalizer.py index d1fad043680..655527269cb 100644 --- a/src/google/adk/flows/llm_flows/core/_finalizer.py +++ b/src/google/adk/flows/llm_flows/core/_finalizer.py @@ -215,8 +215,9 @@ async def handle_after_model_callback( ) -> Optional[LlmResponse]: """Runs after-model callbacks (plugins then agent callbacks). - Also handles grounding metadata injection when google_search_agent is - among the agent's tools. + Also handles grounding metadata injection when a tool sets + ``propagate_grounding_metadata`` and ``temp:_adk_grounding_metadata`` + is present on the session. Args: invocation_context: The invocation context. @@ -238,7 +239,9 @@ async def _maybe_add_grounding_metadata( tools = await agent.canonical_tools(readonly_context) invocation_context.canonical_tools_cache = tools - if not any(tool.name == 'google_search_agent' for tool in tools): + if not any( + getattr(tool, 'propagate_grounding_metadata', False) for tool in tools + ): return response ground_metadata = invocation_context.session.state.get( 'temp:_adk_grounding_metadata', None diff --git a/src/google/adk/tools/mcp_tool/mcp_tool.py b/src/google/adk/tools/mcp_tool/mcp_tool.py index 773637d1d67..3399792f49b 100644 --- a/src/google/adk/tools/mcp_tool/mcp_tool.py +++ b/src/google/adk/tools/mcp_tool/mcp_tool.py @@ -28,7 +28,9 @@ from fastapi.openapi.models import APIKeyIn from google.genai.types import FunctionDeclaration +from google.genai.types import GroundingMetadata from opentelemetry import propagate +from pydantic import ValidationError from typing_extensions import override from ...agents.callback_context import CallbackContext @@ -299,6 +301,7 @@ def __init__( | None ) = None, progress_callback: ProgressFnT | ProgressCallbackFactory | None = None, + propagate_grounding_metadata: bool = False, ): """Initializes an McpTool. @@ -325,6 +328,10 @@ def __init__( The factory receives (tool_name, callback_context, **kwargs) and returns a ProgressFnT or None. This allows callbacks to access and modify runtime context like session state. + propagate_grounding_metadata: If True, copy + ``meta.adk_grounding_metadata`` from the MCP result into + ``temp:_adk_grounding_metadata`` so the flow can attach it to + ``LlmResponse``. Default False. Raises: ValueError: If the MCP tool name collides with a reserved ADK tool @@ -350,6 +357,7 @@ def __init__( self._require_confirmation = require_confirmation self._header_provider = header_provider self._progress_callback = progress_callback + self.propagate_grounding_metadata = propagate_grounding_metadata @override def _get_declaration(self) -> FunctionDeclaration: @@ -724,6 +732,7 @@ async def _run_async_impl( # Keep the caller's key names off the installed SDK's field naming. result = _dump_mcp_model(response) + self._store_grounding_metadata_from_result(result, tool_context) # 2.x-only field. Acting on it (`input_required` drives elicitation) is a # feature, not compatibility. Not dropped on 1.x, where a key of that name @@ -754,6 +763,29 @@ async def _run_async_impl( ) return result + def _store_grounding_metadata_from_result( + self, result: dict[str, Any], tool_context: ToolContext + ) -> None: + """Copies ADK grounding from MCP meta into session temp state.""" + if not self.propagate_grounding_metadata: + return + meta = result.get("meta") + if not isinstance(meta, dict): + return + raw = meta.get("adk_grounding_metadata") + if raw is None: + return + try: + metadata = GroundingMetadata.model_validate(raw) + except ValidationError as e: + logger.warning( + "Ignoring _meta.adk_grounding_metadata from %s: %s", + self.name, + e, + ) + return + tool_context.state["temp:_adk_grounding_metadata"] = metadata + def _detect_error_in_response(self, response: Any) -> str | None: """Telemetry hook: returns an error type if the response indicates an error.""" # `response` is a dumped CallToolResult. `_run_async_impl` restores diff --git a/src/google/adk/tools/mcp_tool/mcp_toolset.py b/src/google/adk/tools/mcp_tool/mcp_toolset.py index 11b9be3fcb3..ed8d5217e6f 100644 --- a/src/google/adk/tools/mcp_tool/mcp_toolset.py +++ b/src/google/adk/tools/mcp_tool/mcp_toolset.py @@ -172,6 +172,7 @@ def __init__( sampling_capabilities: SamplingCapability | None = None, elicitation_callback: ElicitationFnT | None = None, credential_key: str | None = None, + propagate_grounding_metadata: bool = False, ): """Initializes the McpToolset. @@ -224,6 +225,9 @@ def __init__( elicitations used for out-of-band flows such as auth challenges. credential_key: A user specified key used to load and save this credential in a credential service. Used with auth_scheme. + propagate_grounding_metadata: If True, each listed tool copies + ``meta.adk_grounding_metadata`` from the MCP result into + ``temp:_adk_grounding_metadata``. Default False. """ super().__init__(tool_filter=tool_filter, tool_name_prefix=tool_name_prefix) @@ -265,6 +269,7 @@ def __init__( self._auth_scheme = auth_scheme self._auth_credential = auth_credential self._require_confirmation = require_confirmation + self._propagate_grounding_metadata = propagate_grounding_metadata # Store auth config as instance variable so ADK can populate # exchanged_auth_credential in-place before calling get_tools() self._auth_config: Optional[AuthConfig] = ( @@ -540,6 +545,7 @@ async def get_tools( progress_callback=self._progress_callback if hasattr(self, "_progress_callback") else None, + propagate_grounding_metadata=self._propagate_grounding_metadata, ) if self._is_tool_selected(mcp_tool, readonly_context): diff --git a/tests/unittests/flows/llm_flows/core/test_finalizer.py b/tests/unittests/flows/llm_flows/core/test_finalizer.py index c3f249454b3..511547df24c 100644 --- a/tests/unittests/flows/llm_flows/core/test_finalizer.py +++ b/tests/unittests/flows/llm_flows/core/test_finalizer.py @@ -250,6 +250,9 @@ async def test_handle_after_model_callback_grounding_with_callback_override( agent_response.grounding_metadata = state_metadata assert result == agent_response + assert result.grounding_metadata == ( + state_metadata if expect_metadata else None + ) agent_callback.assert_called_once() @@ -311,6 +314,9 @@ def __init__(self): plugin_response.grounding_metadata = state_metadata assert result == plugin_response + assert result.grounding_metadata == ( + state_metadata if expect_metadata else None + ) plugin.after_model_callback.assert_called_once() @@ -324,16 +330,16 @@ async def mock_canonical_tools(self, readonly_context=None): canonical_tools_call_count += 1 from google.adk.tools.base_tool import BaseTool - class MockGoogleSearchTool(BaseTool): + class MockResearchTool(BaseTool): def __init__(self): - super().__init__(name="google_search_agent", description="Mock search") + super().__init__(name="research_agent", description="Mock research") self.propagate_grounding_metadata = True async def call(self, **kwargs): return "mock result" - return [MockGoogleSearchTool()] + return [MockResearchTool()] agent = Agent(name="test_agent", tools=[google_search, dummy_tool]) @@ -376,10 +382,7 @@ async def call(self, **kwargs): assert invocation_context.canonical_tools_cache is not None assert len(invocation_context.canonical_tools_cache) == 1 - assert ( - invocation_context.canonical_tools_cache[0].name - == "google_search_agent" - ) + assert invocation_context.canonical_tools_cache[0].name == "research_agent" assert result1.grounding_metadata == {"foo": "bar"} assert result2.grounding_metadata == {"foo": "bar"} diff --git a/tests/unittests/tools/mcp_tool/test_mcp_tool.py b/tests/unittests/tools/mcp_tool/test_mcp_tool.py index e67e371d2f7..0e63b80e293 100644 --- a/tests/unittests/tools/mcp_tool/test_mcp_tool.py +++ b/tests/unittests/tools/mcp_tool/test_mcp_tool.py @@ -23,6 +23,8 @@ from unittest.mock import patch from google.adk.agents.context import Context +from google.adk.agents.invocation_context import InvocationContext +from google.adk.agents.llm_agent import Agent from google.adk.auth.auth_credential import AuthCredential from google.adk.auth.auth_credential import AuthCredentialTypes from google.adk.auth.auth_credential import HttpAuth @@ -36,6 +38,7 @@ from google.adk.features._feature_registry import temporary_feature_override from google.adk.flows.llm_flows.context import _fencing from google.adk.models.llm_request import LlmRequest +from google.adk.sessions.in_memory_session_service import InMemorySessionService from google.adk.tools.mcp_tool import mcp_tool from google.adk.tools.mcp_tool.mcp_session_manager import _SESSION_IDLE_TTL_SECONDS from google.adk.tools.mcp_tool.mcp_session_manager import MCPSessionManager @@ -45,6 +48,7 @@ from google.adk.tools.mcp_tool.mcp_tool import ProgressFnT from google.adk.tools.tool_context import ToolContext from google.genai.types import FunctionDeclaration +from google.genai.types import GroundingMetadata from mcp.types import CallToolResult from mcp.types import ImageContent from mcp.types import TextContent @@ -789,6 +793,71 @@ async def test_run_async_impl_no_auth(self): "test_tool", arguments=args, progress_callback=None, meta=None ) + async def _tool_context_with_session(self) -> ToolContext: + session_service = InMemorySessionService() + session = await session_service.create_session( + app_name="test_app", user_id="test_user" + ) + tool_context = ToolContext( + invocation_context=InvocationContext( + invocation_id="invocation_id", + agent=Agent(name="test_agent"), + session=session, + session_service=session_service, + ) + ) + tool_context.function_call_id = "test-call-id" + return tool_context + + @pytest.mark.asyncio + async def test_run_async_impl_propagates_grounding_metadata_from_meta(self): + """_meta.adk_grounding_metadata becomes temp state when the flag is on.""" + tool = MCPTool( + mcp_tool=self.mock_mcp_tool, + mcp_session_manager=self.mock_session_manager, + propagate_grounding_metadata=True, + ) + mcp_response = CallToolResult( + content=[TextContent(type="text", text="success")], + _meta={"adk_grounding_metadata": {"webSearchQueries": ["q1"]}}, + ) + self.mock_session.call_tool = AsyncMock(return_value=mcp_response) + tool_context = await self._tool_context_with_session() + + result = await tool._run_async_impl( + args={"param1": "test_value"}, + tool_context=tool_context, + credential=None, + ) + + assert result == expected_tool_result(mcp_response) + stored = tool_context.state["temp:_adk_grounding_metadata"] + assert isinstance(stored, GroundingMetadata) + assert stored.web_search_queries == ["q1"] + + @pytest.mark.asyncio + async def test_run_async_impl_skips_grounding_metadata_when_flag_off(self): + """Default McpTool leaves temp grounding unset even if _meta carries it.""" + tool = MCPTool( + mcp_tool=self.mock_mcp_tool, + mcp_session_manager=self.mock_session_manager, + ) + mcp_response = CallToolResult( + content=[TextContent(type="text", text="success")], + _meta={"adk_grounding_metadata": {"webSearchQueries": ["q1"]}}, + ) + self.mock_session.call_tool = AsyncMock(return_value=mcp_response) + tool_context = await self._tool_context_with_session() + + result = await tool._run_async_impl( + args={"param1": "test_value"}, + tool_context=tool_context, + credential=None, + ) + + assert result == expected_tool_result(mcp_response) + assert "temp:_adk_grounding_metadata" not in tool_context.state + @pytest.mark.asyncio async def test_in_flight_tool_call_is_held_out_of_the_idle_sweep(self): """A call in flight must not have its session swept out from under it.""" diff --git a/tests/unittests/tools/mcp_tool/test_mcp_toolset.py b/tests/unittests/tools/mcp_tool/test_mcp_toolset.py index 4dbfdf6c681..3900db0f2e6 100644 --- a/tests/unittests/tools/mcp_tool/test_mcp_toolset.py +++ b/tests/unittests/tools/mcp_tool/test_mcp_toolset.py @@ -724,6 +724,28 @@ async def my_progress_callback( for tool in tools: assert tool._progress_callback == my_progress_callback + @pytest.mark.asyncio + async def test_get_tools_passes_propagate_grounding_metadata_to_mcp_tools( + self, + ): + """Test that get_tools passes propagate_grounding_metadata to created MCPTool instances.""" + mock_tools = [MockMCPTool("tool1"), MockMCPTool("tool2")] + self.mock_session.list_tools = AsyncMock( + return_value=MockListToolsResult(mock_tools) + ) + + toolset = McpToolset( + connection_params=self.mock_stdio_params, + propagate_grounding_metadata=True, + ) + toolset._mcp_session_manager = self.mock_session_manager + + tools = await toolset.get_tools() + + assert len(tools) == 2 + for tool in tools: + assert tool.propagate_grounding_metadata is True + def test_init_with_progress_callback_factory(self): """Test initialization with a ProgressCallbackFactory.""" From 96319fc85d01c64137610de85250f950055e9747 Mon Sep 17 00:00:00 2001 From: Arjun Ganesh Date: Tue, 29 Sep 2026 15:46:53 -0700 Subject: [PATCH 21/29] docs: explain local lockfile needed by tox in adk-setup skill Merge https://github.com/google/adk-python/pull/7318 PiperOrigin-RevId: 990578226 --- .agents/skills/adk-setup/SKILL.md | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/.agents/skills/adk-setup/SKILL.md b/.agents/skills/adk-setup/SKILL.md index 494aa48ab0e..cbedd8f686d 100644 --- a/.agents/skills/adk-setup/SKILL.md +++ b/.agents/skills/adk-setup/SKILL.md @@ -28,8 +28,9 @@ every dependency extra, pre-commit hooks, and a green unit-test run. python3 --version ``` -2. **uv.** Dependencies are pinned in `uv.lock`; a hand-rolled `pip`/`venv` - environment will not reproduce the locked versions. +2. **uv.** Dependencies are declared in `pyproject.toml`. The `uv sync` step + below creates a local `uv.lock`, which this repository ignores. Run that step + before `tox`, whose lock runner requires the file. ```bash uv --version From 4d06641e7dff5f2f2fa77feea24f68885110cf6d Mon Sep 17 00:00:00 2001 From: Vishal Bulbule Date: Tue, 29 Sep 2026 16:04:16 -0700 Subject: [PATCH 22/29] feat: let BigQuery tools run where CMEK is required Merge https://github.com/google/adk-python/pull/7232 Fixes #3931 PiperOrigin-RevId: 990587308 --- .../bigquery/bigquery_toolset/index.md | 12 ++- .../adk/integrations/bigquery/config.py | 36 +++++++ .../adk/integrations/bigquery/query_tool.py | 26 ++++- .../bigquery/test_bigquery_query_tool.py | 101 ++++++++++++++++++ .../bigquery/test_bigquery_tool_config.py | 24 +++++ 5 files changed, 194 insertions(+), 5 deletions(-) diff --git a/docs/guides/integrations/bigquery/bigquery_toolset/index.md b/docs/guides/integrations/bigquery/bigquery_toolset/index.md index 28ffc7bc888..68c888a6ee6 100644 --- a/docs/guides/integrations/bigquery/bigquery_toolset/index.md +++ b/docs/guides/integrations/bigquery/bigquery_toolset/index.md @@ -74,9 +74,10 @@ Platform. If it is not provided, the tools attempt to use environment-specific defaults. The `bigquery_tool_config` controls the operational limits of the tools. For -example, it defines the maximum number of rows a query can return and whether -the agent is allowed to perform write operations. If this is omitted, the -toolset uses a default `BigQueryToolConfig` instance. +example, it defines the maximum number of rows a query can return, +customer-managed encryption keys (`kms_key_name`), and whether the agent is +allowed to perform write operations. If this is omitted, the toolset uses a +default `BigQueryToolConfig` instance. ## Advanced applications @@ -108,6 +109,11 @@ modules, such as metadata inspection and SQL execution. It does not support every BigQuery API feature, such as managing IAM policies or creating reservation slots. +The `kms_key_name` option on `BigQueryToolConfig` covers `SELECT` results only. +BigQuery rejects a job-level key for DDL, DML, and multi-statement scripts, so +those run without it, requiring a project default key under policies like +`constraints/gcp.restrictNonCmekServices`. + ## Related samples - [bigquery_agent](../../../../../contributing/samples/a2a/a2a_auth/remote_a2a/bigquery_agent/agent.py) - An agent that manages user data on BigQuery using OAuth2. diff --git a/src/google/adk/integrations/bigquery/config.py b/src/google/adk/integrations/bigquery/config.py index cfbbc684ef5..f10d1be90cb 100644 --- a/src/google/adk/integrations/bigquery/config.py +++ b/src/google/adk/integrations/bigquery/config.py @@ -15,6 +15,7 @@ from __future__ import annotations from enum import Enum +import re from typing import Optional from pydantic import BaseModel @@ -148,6 +149,27 @@ class BigQueryToolConfig(BaseModel): "adk-bigquery-" are reserved for internal usage. """ + kms_key_name: Optional[str] = None + """Cloud KMS key to encrypt query results with (CMEK). + + Set this when an organization policy such as + `constraints/gcp.restrictNonCmekServices` requires BigQuery query results to + be protected with a customer-managed key. The value is the key's resource + name, `projects/{project}/locations/{location}/keyRings/{key_ring}/cryptoKeys/{key}`, + and the key must be in the same location as the data being queried. The + BigQuery service agent of the project that runs the query needs the Cloud + KMS CryptoKey Encrypter/Decrypter role on the key. + + The key is applied to SELECT statements only, because BigQuery rejects a + job-level key for DDL, DML, and multi-statement scripts. Under such a policy, + those need a project default key. A permanent table can instead take + `OPTIONS(kms_key_name=...)` in its CREATE statement, but a temporary table + cannot. With `WriteMode.ALLOWED`, setting this adds a dry run before + each query to find the statement type; the other write modes already dry + run the query. For all key options, see + https://cloud.google.com/bigquery/docs/customer-managed-encryption. + """ + @field_validator('maximum_bytes_billed') @classmethod def validate_maximum_bytes_billed(cls, v: Optional[int]) -> Optional[int]: @@ -169,6 +191,20 @@ def validate_application_name(cls, v: Optional[str]) -> Optional[str]: raise ValueError('Application name should not contain spaces.') return v + @field_validator('kms_key_name') + @classmethod + def validate_kms_key_name(cls, v: Optional[str]) -> Optional[str]: + """Validate the Cloud KMS key resource name.""" + if v is not None and not re.fullmatch( + r'projects/[^/]+/locations/[^/]+/keyRings/[^/]+/cryptoKeys/[^/]+', v + ): + raise ValueError( + 'kms_key_name must be a Cloud KMS key resource name of the form' + ' projects/{project}/locations/{location}/keyRings/{key_ring}' + f'/cryptoKeys/{{key}}, found "{v}".' + ) + return v + @field_validator('job_labels') @classmethod def validate_job_labels( diff --git a/src/google/adk/integrations/bigquery/query_tool.py b/src/google/adk/integrations/bigquery/query_tool.py index 9883bed98b6..8c3327e8c3a 100644 --- a/src/google/adk/integrations/bigquery/query_tool.py +++ b/src/google/adk/integrations/bigquery/query_tool.py @@ -211,6 +211,9 @@ def _execute_sql( if settings and settings.application_name: bq_job_labels["adk-bigquery-application-name"] = settings.application_name + # Statement type from a dry run, when the write mode needs one anyway + statement_type: Optional[str] = None + if not settings or settings.write_mode == WriteMode.BLOCKED: dry_run_query_job = bq_client.query( query, @@ -219,7 +222,8 @@ def _execute_sql( dry_run=True, labels=bq_job_labels ), ) - if dry_run_query_job.statement_type != "SELECT": + statement_type = dry_run_query_job.statement_type + if statement_type != "SELECT": return { "status": "ERROR", "error_details": "Read-only mode only supports SELECT statements.", @@ -274,7 +278,8 @@ def _execute_sql( ), ) # A write runs only where the dry run places it in the session dataset. - if dry_run_query_job.statement_type != "SELECT" and not ( + statement_type = dry_run_query_job.statement_type + if statement_type != "SELECT" and not ( dry_run_query_job.destination and dry_run_query_job.destination.dataset_id == bq_session_dataset_id ): @@ -306,6 +311,23 @@ def _execute_sql( ) if settings.maximum_bytes_billed: job_config.maximum_bytes_billed = settings.maximum_bytes_billed + if settings.kms_key_name: + if statement_type is None: + statement_type = bq_client.query( + query, + project=project_id, + job_config=bigquery.QueryJobConfig( + dry_run=True, + connection_properties=bq_connection_properties, + labels=bq_job_labels, + ), + ).statement_type + # BigQuery rejects a job-level key for DDL, DML and scripts, so only the + # results of a SELECT are encrypted with it. + if statement_type == "SELECT": + job_config.destination_encryption_configuration = ( + bigquery.EncryptionConfiguration(kms_key_name=settings.kms_key_name) + ) row_iterator = bq_client.query_and_wait( query, job_config=job_config, diff --git a/tests/unittests/integrations/bigquery/test_bigquery_query_tool.py b/tests/unittests/integrations/bigquery/test_bigquery_query_tool.py index 10d3cda5f40..529ffc2e368 100644 --- a/tests/unittests/integrations/bigquery/test_bigquery_query_tool.py +++ b/tests/unittests/integrations/bigquery/test_bigquery_query_tool.py @@ -2635,6 +2635,107 @@ def test_execute_sql_maximum_bytes_billed_config(): assert call_args.kwargs["job_config"].maximum_bytes_billed == 11_000_000 +_KMS_KEY_NAME = "projects/p/locations/us/keyRings/r/cryptoKeys/k" + + +@pytest.mark.parametrize( + ("write_mode", "query_call_count"), + [ + pytest.param(WriteMode.BLOCKED, 1, id="write-blocked"), + pytest.param(WriteMode.PROTECTED, 2, id="write-protected"), + pytest.param(WriteMode.ALLOWED, 1, id="write-allowed"), + ], +) +def test_execute_sql_encrypts_select_results_with_kms_key( + write_mode, query_call_count +): + """A SELECT runs with the configured KMS key as its destination key. + + Blocked and protected write modes reuse the dry run they already make to + find the statement type. Allowed write mode makes one dry run for it. + """ + credentials = mock.create_autospec(Credentials, instance=True) + tool_config = BigQueryToolConfig( + write_mode=write_mode, kms_key_name=_KMS_KEY_NAME + ) + tool_context = mock.create_autospec(ToolContext, instance=True) + tool_context.state.get.return_value = None + + with mock.patch.object(bigquery, "Client", autospec=True) as Client: + bq_client = Client.return_value + query_job = mock.create_autospec(bigquery.QueryJob) + query_job.statement_type = "SELECT" + bq_client.query.return_value = query_job + + result = query_tool.execute_sql( + "my_project", + "SELECT 123 AS num", + credentials, + tool_config, + tool_context, + ) + + assert result["status"] == "SUCCESS" + assert bq_client.query.call_count == query_call_count + job_config = bq_client.query_and_wait.call_args.kwargs["job_config"] + assert ( + job_config.destination_encryption_configuration.kms_key_name + == _KMS_KEY_NAME + ) + + +def test_execute_sql_does_not_set_kms_key_for_non_select(): + """A statement other than SELECT runs without a job-level KMS key. + + BigQuery rejects a job-level key for DDL, DML and scripts. + """ + credentials = mock.create_autospec(Credentials, instance=True) + tool_config = BigQueryToolConfig( + write_mode=WriteMode.ALLOWED, kms_key_name=_KMS_KEY_NAME + ) + tool_context = mock.create_autospec(ToolContext, instance=True) + + with mock.patch.object(bigquery, "Client", autospec=True) as Client: + bq_client = Client.return_value + query_job = mock.create_autospec(bigquery.QueryJob) + query_job.statement_type = "CREATE_TABLE" + bq_client.query.return_value = query_job + + result = query_tool.execute_sql( + "my_project", + "CREATE TABLE ds.t AS SELECT 1 AS x", + credentials, + tool_config, + tool_context, + ) + + assert result["status"] == "SUCCESS" + job_config = bq_client.query_and_wait.call_args.kwargs["job_config"] + assert job_config.destination_encryption_configuration is None + + +def test_execute_sql_without_kms_key_adds_no_dry_run(): + """Without a KMS key, allowed write mode still makes no dry run.""" + credentials = mock.create_autospec(Credentials, instance=True) + tool_config = BigQueryToolConfig(write_mode=WriteMode.ALLOWED) + tool_context = mock.create_autospec(ToolContext, instance=True) + + with mock.patch.object(bigquery, "Client", autospec=True) as Client: + bq_client = Client.return_value + + query_tool.execute_sql( + "my_project", + "SELECT 123 AS num", + credentials, + tool_config, + tool_context, + ) + + bq_client.query.assert_not_called() + job_config = bq_client.query_and_wait.call_args.kwargs["job_config"] + assert job_config.destination_encryption_configuration is None + + @pytest.mark.parametrize( ("tool_call",), [ diff --git a/tests/unittests/integrations/bigquery/test_bigquery_tool_config.py b/tests/unittests/integrations/bigquery/test_bigquery_tool_config.py index d81b5d83e92..5abc5f7bbb7 100644 --- a/tests/unittests/integrations/bigquery/test_bigquery_tool_config.py +++ b/tests/unittests/integrations/bigquery/test_bigquery_tool_config.py @@ -76,6 +76,30 @@ def test_bigquery_tool_config_invalid_maximum_bytes_billed(): BigQueryToolConfig(maximum_bytes_billed=10_485_759) +def test_bigquery_tool_config_valid_kms_key_name(): + """Test BigQueryToolConfig accepts a Cloud KMS key resource name.""" + key = "projects/p/locations/us/keyRings/r/cryptoKeys/k" + config = BigQueryToolConfig(kms_key_name=key) + assert config.kms_key_name == key + + +@pytest.mark.parametrize( + "key", + [ + pytest.param("k", id="bare-key-id"), + pytest.param( + "projects/p/locations/us/keyRings/r/cryptoKeys/k/cryptoKeyVersions/1", + id="key-version", + ), + pytest.param("projects/p/locations/us/keyRings/r", id="key-ring"), + ], +) +def test_bigquery_tool_config_invalid_kms_key_name(key): + """Test BigQueryToolConfig rejects a value that is not a key resource name.""" + with pytest.raises(ValueError, match="kms_key_name must be a Cloud KMS key"): + BigQueryToolConfig(kms_key_name=key) + + @pytest.mark.parametrize( "labels", [ From b2da4c6030dcd4b4f156743201a49f01bd4ec61a Mon Sep 17 00:00:00 2001 From: Xuan Yang Date: Tue, 29 Sep 2026 16:51:52 -0700 Subject: [PATCH 23/29] feat: add the model consult session context handover layer When a small executor model escalates a hard step to a larger advisor model, the advisor has to be told what has happened so far, or it answers a question it does not understand. This adds the layer that turns a session event log into contents the advisor can read: - tool calls and tool results are flattened into plain text, so a model that was never given those tool declarations can still follow them; - consecutive events from the same role are merged into one turn, which third-party advisor models require; - the transcript is held to a character budget by keeping the head and the tail, marking the omitted middle, and trimming the newest turn when it does not fit on its own, so one long turn cannot silently multiply the cost of a consult; - thoughts, media parts, per-part length and event count are each configurable, and the in-flight escalation call itself is skipped. The layer is not wired into a tool yet; that follows in a later change. Co-authored-by: Xuan Yang PiperOrigin-RevId: 990610529 --- .../adk/tools/model_consult/__init__.py | 21 + .../adk/tools/model_consult/_context.py | 475 +++++++++++++++++ .../unittests/tools/model_consult/__init__.py | 13 + .../tools/model_consult/test_context.py | 483 ++++++++++++++++++ 4 files changed, 992 insertions(+) create mode 100644 src/google/adk/tools/model_consult/__init__.py create mode 100644 src/google/adk/tools/model_consult/_context.py create mode 100644 tests/unittests/tools/model_consult/__init__.py create mode 100644 tests/unittests/tools/model_consult/test_context.py diff --git a/src/google/adk/tools/model_consult/__init__.py b/src/google/adk/tools/model_consult/__init__.py new file mode 100644 index 00000000000..a469d29c802 --- /dev/null +++ b/src/google/adk/tools/model_consult/__init__.py @@ -0,0 +1,21 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Lets a fast executor model consult a stronger advisor model mid-task.""" + +from ._context import ModelConsultContextConfig + +__all__ = [ + 'ModelConsultContextConfig', +] diff --git a/src/google/adk/tools/model_consult/_context.py b/src/google/adk/tools/model_consult/_context.py new file mode 100644 index 00000000000..75f7a58496c --- /dev/null +++ b/src/google/adk/tools/model_consult/_context.py @@ -0,0 +1,475 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Hands the executor's session over to the advisor model. + +What makes a consult worth more than a plain call to a larger model is that +the advisor sees what the executor saw: the same instructions, the same tool +results, the same dead ends. This module turns a session's event log into +content any advisor model can read, under a character budget so that a +long-horizon session cannot silently blow up the cost of a single consult. +""" + +from __future__ import annotations + +from collections.abc import Sequence +import json +from typing import Any +from typing import TYPE_CHECKING + +from google.genai import types +from pydantic import BaseModel +from pydantic import ConfigDict +from pydantic import Field + +from ...events._rewind_events import _apply_rewinds + +if TYPE_CHECKING: + from ...events.event import Event + +_ROLE_LABELS = {'user': 'USER', 'model': 'AGENT'} +# How much more room plain text gets than a rendered tool payload. Documented +# on ModelConsultContextConfig.max_part_chars. +_TEXT_CHARS_MULTIPLIER = 8 +_OMISSION_MARKER = ( + '[... {n} earlier turn(s) omitted to fit the context budget ...]' +) + + +class ModelConsultContextConfig(BaseModel): + """Controls how much of the executor's session reaches the advisor.""" + + model_config = ConfigDict(extra='forbid') + + include_session: bool = True + """Whether to send the session at all. + + When False, the advisor only sees the question and context that the executor + passed as tool arguments. + """ + + max_events: int | None = Field(default=None, ge=1) + """Keep at most this many of the most recent events. None keeps all. + + Counted over raw session events, before thoughts and other withheld parts + are filtered out, so the advisor may end up seeing fewer turns than this. + """ + + max_chars: int | None = Field(default=200_000, ge=1) + """Character budget for the handover. + + Whole turns are dropped from the middle once the budget is exceeded: the + original task and the most recent turns are what the advisor needs. The + newest turn is always kept, trimmed if it does not fit on its own. + """ + + max_part_chars: int = Field(default=4_000, ge=1) + """Per-part cap on rendered tool calls and tool results. + + These are the usual source of runaway context. Plain model text gets + `_TEXT_CHARS_MULTIPLIER` times this allowance, since prose is rarely what + blows a session up and cutting an answer mid-sentence costs the advisor + more than it saves. + """ + + include_media: bool = True + """Whether to pass inline images and audio through to the advisor. + + Turn this off for text-only advisor models. + """ + + include_thoughts: bool = False + """Whether to include the executor's own thought parts. + + Off by default: thought summaries are noisy, and they bias the advisor + toward the framing the executor is already stuck in. + """ + + +def _truncate(text: str, limit: int) -> str: + """Cuts `text` down to `limit` characters, noting how much was dropped. + + The note counts against the limit: a cap the caller set is a cap on what + actually gets sent, not on what is left after the note is added. + + Args: + text: The text to cut down. + limit: The character budget for the returned string. + + Returns: + The text, at most `limit` characters long. + """ + if limit <= 0: + return '' + if len(text) <= limit: + return text + # Sized against the largest count that could be reported, so that the note + # never pushes the result back over the limit. + widest_note = f'\n[... {len(text)} characters truncated ...]' + keep = limit - len(widest_note) + if keep <= 0: + return text[:limit] + return f'{text[:keep]}\n[... {len(text) - keep} characters truncated ...]' + + +def _render_args(args: dict[str, Any] | None, limit: int) -> str: + """Renders function call arguments as a JSON string.""" + if not args: + return '' + try: + rendered = json.dumps(args, ensure_ascii=False, default=str) + except (TypeError, ValueError): + rendered = str(args) + return _truncate(rendered, limit) + + +def _render_response(response: Any, limit: int) -> str: + """Renders a function response body as a string.""" + if response is None: + return '' + if isinstance(response, str): + rendered = response + else: + try: + rendered = json.dumps(response, ensure_ascii=False, default=str) + except (TypeError, ValueError): + rendered = str(response) + return _truncate(rendered, limit) + + +def _convert_part( + part: types.Part, + config: ModelConsultContextConfig, + skip_function_call_ids: frozenset[str], +) -> types.Part | None: + """Normalizes one part into something any advisor model can read. + + Function calls and responses become readable text rather than live tool + parts: the advisor does not hold the executor's tool declarations, and a + dangling function call is a validation error for most providers. + + Args: + part: The part to convert. + config: The handover configuration. + skip_function_call_ids: Function call ids to drop entirely. + + Returns: + The converted part, or None when the part carries nothing worth sending. + """ + if part.thought and not config.include_thoughts: + return None + + if part.function_call is not None: + call = part.function_call + if call.id and call.id in skip_function_call_ids: + return None + args = _render_args(call.args, config.max_part_chars) + return types.Part(text=f'[tool_call] {call.name}({args})') + + if part.function_response is not None: + response = part.function_response + if response.id and response.id in skip_function_call_ids: + return None + body = _render_response(response.response, config.max_part_chars) + return types.Part(text=f'[tool_result] {response.name} -> {body}') + + if part.text is not None: + text = _truncate(part.text, config.max_part_chars * _TEXT_CHARS_MULTIPLIER) + if not text.strip(): + return None + if part.thought: + # Labelled, because rebuilding the part drops `thought` and the advisor + # is being asked to doubt exactly this reasoning: it has to be able to + # tell it apart from what the executor actually concluded. + text = f'[thought] {text}' + return types.Part(text=text) + + if part.inline_data is not None or part.file_data is not None: + if config.include_media: + return part + description = _describe_media_part(part, reason='omitted') + return None if description is None else types.Part(text=description) + + if part.executable_code is not None: + code = _truncate(part.executable_code.code or '', config.max_part_chars) + return types.Part(text=f'[code]\n{code}') + + if part.code_execution_result is not None: + output = _truncate( + part.code_execution_result.output or '', config.max_part_chars + ) + return types.Part(text=f'[code_result] {output}') + + return None + + +def _part_chars(part: types.Part) -> int: + """Estimates how much of the character budget one part consumes.""" + if part.text: + return len(part.text) + if part.inline_data is not None and part.inline_data.data: + # Rough stand-in so that media still consumes budget. + return len(part.inline_data.data) // 4 + return 0 + + +def _content_chars(content: types.Content) -> int: + """Estimates how much of the character budget one content consumes.""" + return sum(_part_chars(part) for part in content.parts or []) + + +def _merge_adjacent(contents: Sequence[types.Content]) -> list[types.Content]: + """Collapses consecutive same-role contents into one. + + Gemini tolerates consecutive user turns, but several third-party advisor + models reached through LiteLlm require strict role alternation, so the + handover is normalized before it leaves. + + Args: + contents: The contents to normalize, in order. + + Returns: + The contents with adjacent same-role entries merged. + """ + merged: list[types.Content] = [] + for content in contents: + if merged and merged[-1].role == content.role: + merged[-1] = types.Content( + role=content.role, + parts=list(merged[-1].parts or []) + list(content.parts or []), + ) + else: + merged.append(content) + return merged + + +def _truncate_content(content: types.Content, limit: int) -> types.Content: + """Trims a content down to `limit` characters. + + Media is charged against the budget on the same rough basis the budget was + measured with, and replaced by a placeholder when it does not fit: a turn + made of images would otherwise sail past the cap untouched. + + Args: + content: The content to trim. + limit: The character budget for this content. + + Returns: + The content, trimmed to fit. + """ + parts: list[types.Part] = [] + used = 0 + for part in content.parts or []: + if used >= limit: + continue + if part.text is not None: + parts.append(types.Part(text=_truncate(part.text, limit - used))) + used += len(part.text) + continue + cost = _part_chars(part) + if cost <= limit - used: + parts.append(part) + used += cost + continue + placeholder = _describe_media_part( + part, reason='omitted to fit the context budget' + ) + if placeholder is not None: + placeholder = _truncate(placeholder, limit - used) + if placeholder: + parts.append(types.Part(text=placeholder)) + used += len(placeholder) + return types.Content(role=content.role, parts=parts) + + +def _describe_media_part(part: types.Part, *, reason: str = '') -> str | None: + """Names a non-text part in plain text, so its absence stays visible. + + One describer for every renderer: the transcript, the text-only conversion + and the budget trim all name a part the same way, and a part kind that is + handled here cannot be silently dropped by one of them. + + Args: + part: The part to name. + reason: Why the part is named instead of carried, when it was dropped. + + Returns: + A bracketed description, or None when the part carries no media. + """ + suffix = f' {reason}' if reason else '' + if part.inline_data is not None: + mime_type = part.inline_data.mime_type or 'unknown' + return f'[media{suffix}: {mime_type}]' + if part.file_data is not None: + return f'[file{suffix}: {part.file_data.file_uri}]' + return None + + +def _apply_char_budget( + contents: list[types.Content], max_chars: int | None +) -> list[types.Content]: + """Drops whole turns from the middle until the budget is met. + + Keeping the head preserves the original task; keeping the tail preserves the + state the executor is actually stuck in. The newest turn is always kept, so + when it alone is larger than the budget it is trimmed rather than allowed to + undo the budget. + + Args: + contents: The contents to trim, in order. + max_chars: The character budget, or None for no budget. + + Returns: + The contents, with an omission marker in place of any dropped turns. + """ + if max_chars is None or not contents: + return contents + + sizes = [_content_chars(content) for content in contents] + if sum(sizes) <= max_chars: + return contents + + head_budget = max_chars // 4 + head: list[types.Content] = [] + used = 0 + for content, size in zip(contents, sizes): + if used + size > head_budget: + break + head.append(content) + used += size + + # The marker is part of what gets sent, so it comes out of the budget too. + remaining = max(max_chars - used - len(_OMISSION_MARKER), 0) + tail: list[types.Content] = [] + tail_used = 0 + for content, size in zip( + reversed(contents[len(head) :]), reversed(sizes[len(head) :]) + ): + if tail_used + size > remaining and tail: + break + tail.append(content) + tail_used += size + tail.reverse() + + dropped = len(contents) - len(head) - len(tail) + marker = ( + types.Content( + role='user', + parts=[types.Part(text=_OMISSION_MARKER.format(n=dropped))], + ) + if dropped > 0 + else None + ) + + # The tail loop takes the newest turn whatever its size, so it is the one + # place the budget can still be blown. Trim that turn instead of reporting a + # cap the handover does not honour. + newest = tail[-1] + # Everything kept except the newest turn, which is what is left to trim. + fixed = used + tail_used - _content_chars(newest) + if marker is not None: + if max_chars - fixed - _content_chars(marker) <= 0: + # A budget this small cannot carry both. The newest turn is the state + # the advisor is being asked about, so the marker is what goes. + marker = None + else: + fixed += _content_chars(marker) + + kept = head + ([marker] if marker is not None else []) + tail + allowance = max_chars - fixed + if allowance < _content_chars(newest): + kept = kept[:-1] + [_truncate_content(newest, max(allowance, 0))] + + # A turn that trimmed down to nothing is dropped rather than sent: a content + # with no parts is a validation error for several providers, and it tells the + # advisor nothing anyway. + return [content for content in kept if content.parts] + + +def build_advisor_contents( + events: Sequence[Event], + *, + config: ModelConsultContextConfig | None = None, + skip_function_call_ids: Sequence[str] = (), +) -> list[types.Content]: + """Converts session events into contents for the advisor request. + + Args: + events: The session's event log, oldest first. + config: The handover configuration. Defaults are used when omitted. + skip_function_call_ids: Function call ids to drop, normally the in-flight + consult itself, which the handoff message restates anyway. + + Returns: + Normalized, budget-bounded contents. Empty when there is nothing to send. + """ + config = config or ModelConsultContextConfig() + if not config.include_session: + return [] + + skipped = frozenset(call_id for call_id in skip_function_call_ids if call_id) + # Rewound invocations are still in the log but the executor no longer sees + # them, and the point of the handover is that the advisor sees what the + # executor saw. Same helper the prompt builder and the compactor use. + live = _apply_rewinds(list(events)) + kept = [event for event in live if not event.partial] + if config.max_events is not None: + kept = kept[-config.max_events :] + + contents: list[types.Content] = [] + for event in kept: + content = event.content + if content is None or not content.parts: + continue + parts = [ + converted + for part in content.parts + if (converted := _convert_part(part, config, skipped)) is not None + ] + if not parts: + continue + author = event.author or content.role or 'model' + role = 'user' if author == 'user' else 'model' + contents.append(types.Content(role=role, parts=parts)) + + # Merged twice on purpose: the budget pass can splice an omission marker + # between two turns of the same role, which is exactly what the first merge + # was there to rule out. + trimmed = _apply_char_budget(_merge_adjacent(contents), config.max_chars) + return _merge_adjacent(trimmed) + + +def render_transcript(contents: Sequence[types.Content]) -> str: + """Renders contents as a labelled plain-text transcript. + + Args: + contents: The contents to render, in order. + + Returns: + The rendered transcript, with one labelled block per content. + """ + lines: list[str] = [] + for content in contents: + label = _ROLE_LABELS.get(content.role or 'model', 'AGENT') + chunks: list[str] = [] + for part in content.parts or []: + if part.text: + chunks.append(part.text) + continue + description = _describe_media_part(part) + if description is not None: + chunks.append(description) + if chunks: + lines.append(f'{label}: ' + '\n'.join(chunks)) + return '\n\n'.join(lines) diff --git a/tests/unittests/tools/model_consult/__init__.py b/tests/unittests/tools/model_consult/__init__.py new file mode 100644 index 00000000000..58d482ea386 --- /dev/null +++ b/tests/unittests/tools/model_consult/__init__.py @@ -0,0 +1,13 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. diff --git a/tests/unittests/tools/model_consult/test_context.py b/tests/unittests/tools/model_consult/test_context.py new file mode 100644 index 00000000000..571c2c97603 --- /dev/null +++ b/tests/unittests/tools/model_consult/test_context.py @@ -0,0 +1,483 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Tests for the model consult session-to-advisor handover.""" + +from typing import Sequence + +from google.adk.events.event import Event +from google.adk.tools.model_consult._context import build_advisor_contents +from google.adk.tools.model_consult._context import ModelConsultContextConfig +from google.adk.tools.model_consult._context import render_transcript +from google.genai import types +from pydantic import ValidationError +import pytest + +# ADK authors every event the agent produces, including tool results, with the +# agent's own name. Only the end user's turns are authored 'user'. +_AGENT = 'root_agent' + + +def _user_event(text: str) -> Event: + """Builds a user turn carrying a single text part.""" + return Event( + author='user', + content=types.Content(role='user', parts=[types.Part(text=text)]), + ) + + +def _agent_event(parts: list[types.Part]) -> Event: + """Builds an agent turn carrying the given parts.""" + return Event(author=_AGENT, content=types.Content(role='model', parts=parts)) + + +def _tool_result_event( + name: str, response: dict[str, object], *, call_id: str = 'fc-1' +) -> Event: + """Builds a tool result event the way ADK's tool caller builds it. + + The author is the agent, not the user; only `content.role` is 'user'. See + `flows/llm_flows/tools/_caller.py`, which sets `function_response.id`, builds + the response content with `role='user'`, and authors the event with the + agent's name. + """ + return Event( + author=_AGENT, + content=types.Content( + role='user', + parts=[ + types.Part( + function_response=types.FunctionResponse( + id=call_id, name=name, response=response + ) + ) + ], + ), + ) + + +def _texts(contents: Sequence[types.Content]) -> list[str]: + """Flattens the text of every part, in order.""" + return [ + part.text or '' for content in contents for part in content.parts or [] + ] + + +def _chars(contents: Sequence[types.Content]) -> int: + """Counts the characters the handover would actually send.""" + return sum(len(text) for text in _texts(contents)) + + +def test_session_is_replayed_as_multi_turn_contents(): + """Events reach the advisor in order, with their roles preserved.""" + events = [ + _user_event('Investigate the paging alert.'), + _agent_event([types.Part(text='Checking logs.')]), + ] + + contents = build_advisor_contents(events) + + assert [content.role for content in contents] == ['user', 'model'] + assert _texts(contents) == ['Investigate the paging alert.', 'Checking logs.'] + + +def test_executor_thoughts_are_withheld_by_default(): + """Thought parts do not reach the advisor unless asked for.""" + events = [ + _agent_event([ + types.Part(text='internal musing', thought=True), + types.Part(text='visible answer'), + ]) + ] + + contents = build_advisor_contents(events) + + assert _texts(contents) == ['visible answer'] + + +def test_included_thoughts_are_labelled_as_thoughts(): + """Reasoning stays distinguishable from what the executor concluded.""" + events = [ + _agent_event([ + types.Part(text='internal musing', thought=True), + types.Part(text='visible answer'), + ]) + ] + + contents = build_advisor_contents( + events, config=ModelConsultContextConfig(include_thoughts=True) + ) + + assert _texts(contents) == ['[thought] internal musing', 'visible answer'] + + +def test_tool_calls_and_results_are_flattened_into_text(): + """Function parts become readable text the advisor can consume. + + The advisor holds none of the executor's tool declarations, so a live + function call part would be a validation error for most providers. + """ + events = [ + _agent_event([ + types.Part( + function_call=types.FunctionCall( + id='fc-1', name='query_logs', args={'service': 'checkout'} + ) + ) + ]), + _tool_result_event('query_logs', {'errors': 42}), + ] + + contents = build_advisor_contents(events) + + # The tool result is authored by the agent, so it lands in the same model + # turn as the call that produced it. + assert [content.role for content in contents] == ['model'] + assert _texts(contents) == [ + '[tool_call] query_logs({"service": "checkout"})', + '[tool_result] query_logs -> {"errors": 42}', + ] + assert all( + part.function_call is None and part.function_response is None + for content in contents + for part in content.parts or [] + ) + + +def test_in_flight_consult_is_left_out_of_the_handover(): + """The consult that triggered the handover is not replayed back to it.""" + events = [ + _agent_event([ + types.Part( + function_call=types.FunctionCall( + id='fc-current', + name='model_consult', + args={'question': 'help'}, + ) + ) + ]) + ] + + contents = build_advisor_contents( + events, skip_function_call_ids=['fc-current'] + ) + + assert not contents + + +def test_in_flight_consult_result_is_left_out_of_the_handover(): + """The matching tool result is skipped by the same id.""" + events = [ + _tool_result_event( + 'model_consult', {'status': 'ok'}, call_id='fc-current' + ), + ] + + contents = build_advisor_contents( + events, skip_function_call_ids=['fc-current'] + ) + + assert not contents + + +def test_consecutive_same_role_turns_are_merged(): + """Adjacent same-role turns collapse into one content. + + Advisor models reached through LiteLlm require strict role alternation. + """ + events = [ + _agent_event([types.Part(text='one')]), + _agent_event([types.Part(text='two')]), + ] + + contents = build_advisor_contents(events) + + assert len(contents) == 1 + assert _texts(contents) == ['one', 'two'] + + +def test_partial_streaming_events_are_ignored(): + """Streaming fragments are skipped so text is not duplicated.""" + streaming = _agent_event([types.Part(text='partial chunk')]) + streaming.partial = True + + contents = build_advisor_contents([streaming, _user_event('done')]) + + assert _texts(contents) == ['done'] + + +def test_max_events_keeps_only_the_most_recent_turns(): + """The event cap trims from the front, keeping the newest turns.""" + events = [_user_event(f'turn {i}') for i in range(10)] + + contents = build_advisor_contents( + events, config=ModelConsultContextConfig(max_events=3) + ) + + assert _texts(contents) == ['turn 7', 'turn 8', 'turn 9'] + + +def test_character_budget_drops_the_middle_and_marks_the_gap(): + """Over budget, the original task and the current state both survive.""" + events = [] + for i in range(20): + events.append(_user_event(f'user {i} ' + 'x' * 500)) + events.append(_agent_event([types.Part(text=f'model {i} ' + 'y' * 500)])) + + contents = build_advisor_contents( + events, config=ModelConsultContextConfig(max_chars=4000) + ) + + texts = _texts(contents) + assert any('omitted to fit the context budget' in text for text in texts) + assert texts[0].startswith('user 0') + assert texts[-1].startswith('model 19') + assert _chars(contents) <= 4000 + + +def test_budget_survives_one_turn_larger_than_the_whole_budget(): + """The newest turn is always kept, so it is trimmed rather than exempted.""" + events = [ + _user_event('small task'), + _agent_event([types.Part(text='Z' * 30_000)]), + ] + + contents = build_advisor_contents( + events, config=ModelConsultContextConfig(max_chars=1000) + ) + + assert _chars(contents) <= 1000 + assert 'characters truncated' in _texts(contents)[-1] + + +@pytest.mark.parametrize('max_chars', [40, 1000]) +def test_budget_holds_when_the_newest_turn_is_media(max_chars: int): + """Media and its omission placeholder both count against the budget.""" + events = [ + _user_event('small task'), + _agent_event([ + types.Part( + inline_data=types.Blob(mime_type='image/png', data=b'x' * 40_000) + ) + ]), + ] + + contents = build_advisor_contents( + events, config=ModelConsultContextConfig(max_chars=max_chars) + ) + + assert _chars(contents) <= max_chars + assert _texts(contents)[-1].startswith('[media omitted to fit') + + +@pytest.mark.parametrize('max_chars', [40, 61]) +def test_budget_too_small_for_the_marker_keeps_the_newest_turn(max_chars: int): + """The newest turn outranks the omission marker, and never ships empty. + + A content with no parts is a validation error for several providers, so a + budget that cannot carry both (including `max_chars=61`, the exact length of + the marker itself) has to drop the marker, not the turn. + """ + events = [ + _user_event('a' * 200), + _agent_event([types.Part(text='b' * 200)]), + _user_event('c' * 200), + _agent_event([types.Part(text='d' * 400)]), + ] + + contents = build_advisor_contents( + events, config=ModelConsultContextConfig(max_chars=max_chars) + ) + + assert contents + assert all(content.parts for content in contents) + assert _chars(contents) <= max_chars + assert 'd' in _texts(contents)[-1] + + +def test_rewound_invocations_are_not_handed_over(): + """The executor no longer sees a rewound turn, so neither does the advisor.""" + discarded = _user_event('wrong task') + discarded.invocation_id = 'inv1' + rewind = Event(author='user', invocation_id='inv2') + rewind.actions.rewind_before_invocation_id = 'inv1' + live = _user_event('real task') + live.invocation_id = 'inv3' + + contents = build_advisor_contents([discarded, rewind, live]) + + assert _texts(contents) == ['real task'] + + +def test_trimming_keeps_the_roles_alternating(): + """The omission marker must not re-introduce adjacent same-role turns.""" + events = [] + for i in range(20): + events.append(_user_event(f'user {i} ' + 'x' * 500)) + events.append(_agent_event([types.Part(text=f'model {i} ' + 'y' * 500)])) + + contents = build_advisor_contents( + events, config=ModelConsultContextConfig(max_chars=4000) + ) + + roles = [content.role for content in contents] + assert all(before != after for before, after in zip(roles, roles[1:])) + + +def test_oversized_tool_results_are_truncated_per_part(): + """A single huge tool result cannot consume the whole handover.""" + events = [_tool_result_event('dump', {'blob': 'z' * 50_000})] + + contents = build_advisor_contents( + events, config=ModelConsultContextConfig(max_part_chars=500) + ) + + text = _texts(contents)[0] + assert 'characters truncated' in text + # The 500 characters that survive, plus the prefix and the truncation note. + assert len(text) < 600 + + +def test_plain_text_gets_more_room_than_a_tool_result(): + """Prose receives _TEXT_CHARS_MULTIPLIER times the per-part tool cap.""" + prose = 'p' * 3000 + events = [ + _agent_event([types.Part(text=prose)]), + _tool_result_event('dump', {'blob': 'z' * 3000}), + ] + + contents = build_advisor_contents( + events, config=ModelConsultContextConfig(max_part_chars=500) + ) + + texts = _texts(contents) + assert texts[0] == prose + assert 'characters truncated' in texts[1] + + +def test_media_reaches_the_advisor_by_default(): + """Inline media is passed through untouched.""" + media = types.Part( + inline_data=types.Blob(mime_type='image/png', data=b'\x89PNG fake') + ) + + contents = build_advisor_contents([_agent_event([media])]) + + assert contents[0].parts[0].inline_data is not None + + +def test_media_is_described_in_text_for_text_only_advisors(): + """With include_media off, media becomes a placeholder instead.""" + media = types.Part( + inline_data=types.Blob(mime_type='image/png', data=b'\x89PNG fake') + ) + + contents = build_advisor_contents( + [_agent_event([media])], + config=ModelConsultContextConfig(include_media=False), + ) + + assert _texts(contents) == ['[media omitted: image/png]'] + + +def test_code_parts_are_rendered_as_text(): + """Executed code and its output reach the advisor as readable text.""" + events = [ + _agent_event([ + types.Part( + executable_code=types.ExecutableCode( + code='print(1)', language=types.Language.PYTHON + ) + ), + types.Part( + code_execution_result=types.CodeExecutionResult( + outcome=types.Outcome.OUTCOME_OK, output='1' + ) + ), + ]) + ] + + contents = build_advisor_contents(events) + + assert _texts(contents) == ['[code]\nprint(1)', '[code_result] 1'] + + +def test_whitespace_only_text_is_dropped(): + """Blank turns are not worth a slot in the handover.""" + events = [_agent_event([types.Part(text=' \n ')])] + + contents = build_advisor_contents(events) + + assert not contents + + +def test_session_can_be_withheld_entirely(): + """With include_session off, the advisor sees no session content.""" + events = [_user_event('secret internal transcript')] + + contents = build_advisor_contents( + events, config=ModelConsultContextConfig(include_session=False) + ) + + assert not contents + + +def test_transcript_rendering_labels_each_role(): + """Transcript mode renders contents as a labelled plain-text block.""" + events = [_user_event('question'), _agent_event([types.Part(text='answer')])] + + transcript = render_transcript(build_advisor_contents(events)) + + assert transcript == 'USER: question\n\nAGENT: answer' + + +def test_transcript_rendering_names_media_it_cannot_write_out(): + """Media survives as a marker so the transcript is not silently lossy.""" + media = types.Part( + inline_data=types.Blob(mime_type='image/png', data=b'\x89PNG fake') + ) + + transcript = render_transcript( + build_advisor_contents([_agent_event([media])]) + ) + + assert transcript == 'AGENT: [media: image/png]' + + +def test_transcript_rendering_names_file_parts(): + """A file part carries no text and no bytes, so it is the easiest to lose.""" + file_part = types.Part( + file_data=types.FileData( + file_uri='gs://bucket/spec.pdf', mime_type='application/pdf' + ) + ) + + transcript = render_transcript( + build_advisor_contents([_agent_event([file_part])]) + ) + + assert transcript == 'AGENT: [file: gs://bucket/spec.pdf]' + + +def test_config_rejects_unknown_fields(): + """A misspelled option fails loudly instead of being silently ignored.""" + with pytest.raises(ValidationError): + ModelConsultContextConfig(max_char=100) + + +@pytest.mark.parametrize('field', ['max_events', 'max_chars', 'max_part_chars']) +def test_config_rejects_degenerate_caps(field: str): + """A cap of zero once meant 'no cap', which is the opposite of the ask.""" + with pytest.raises(ValidationError): + ModelConsultContextConfig(**{field: 0}) From 28c47b550d29746ce4f4b38a30617e0f6e14c1eb Mon Sep 17 00:00:00 2001 From: Xuan Yang Date: Tue, 29 Sep 2026 16:58:22 -0700 Subject: [PATCH 24/29] feat(tools): invoke advisor models without tools for model_consult Invoke the advisor `BaseLlm` via `generate_content_async(req, stream=False)` with `config.tools = []` and `config.tool_config = None` so mid-task consultations return text guidance without entering a tool loop, while still emitting standard OpenTelemetry client duration and token usage metrics: - `resolve_advisor_llm` and `resolve_thinking_level` for model and thinking-level normalization - `call_advisor` with automatic fallback retry when `thinking_config` is rejected, `MAX_TOKENS` truncation and thought-exhaustion handling across `Gemini` and `LiteLlm`, and OpenTelemetry metric emission - `AdvisorResult`, `AdvisorUsage`, and `AdvisorError` Co-authored-by: Xuan Yang PiperOrigin-RevId: 990613834 --- .../adk/tools/model_consult/_advisor.py | 551 +++++++++++++ .../tools/model_consult/test_advisor.py | 730 ++++++++++++++++++ 2 files changed, 1281 insertions(+) create mode 100644 src/google/adk/tools/model_consult/_advisor.py create mode 100644 tests/unittests/tools/model_consult/test_advisor.py diff --git a/src/google/adk/tools/model_consult/_advisor.py b/src/google/adk/tools/model_consult/_advisor.py new file mode 100644 index 00000000000..625dac02e22 --- /dev/null +++ b/src/google/adk/tools/model_consult/_advisor.py @@ -0,0 +1,551 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Advisor LLM invocation for the `model_consult` tool. + +Calls a `BaseLlm` directly via `generate_content_async(req, stream=False)` +with `config.tools = []` and `config.tool_config = None` so the advisor returns +text guidance only and cannot call tools or enter an agent loop. +""" + +from __future__ import annotations + +import asyncio +from collections.abc import AsyncGenerator +from collections.abc import Sequence +import copy +from dataclasses import dataclass +import logging +import time + +from google.genai import types + +from ...models.base_llm import BaseLlm +from ...models.llm_request import LlmRequest +from ...models.llm_response import LlmResponse +from ...models.registry import LLMRegistry +from ...telemetry import _metrics +from ...telemetry import tracing +from ...telemetry._token_usage import TokenUsage +from ...utils.model_name_utils import is_gemini_model + +logger = logging.getLogger('google_adk.' + __name__) + +_THINKING_LEVEL_MAP: dict[str, types.ThinkingLevel] = { + 'minimal': types.ThinkingLevel.MINIMAL, + 'low': types.ThinkingLevel.LOW, + 'medium': types.ThinkingLevel.MEDIUM, + 'high': types.ThinkingLevel.HIGH, +} + + +class AdvisorError(RuntimeError): # pylint: disable=g-bad-exception-name + """Raised when the advisor model call fails or returns unusable output.""" + + +@dataclass(frozen=True, kw_only=True) +class AdvisorUsage: + """Token accounting for a single advisor call (or cumulative across calls). + + Attributes: + prompt_tokens: Input tokens billed for the prompt (including tool-use prompt + tokens when reported). + output_tokens: Candidate output tokens (excluding thoughts). + thoughts_tokens: Reasoning/thinking tokens consumed by the advisor. + cached_tokens: Prompt tokens served from a context cache. + total_tokens: Total tokens consumed (`prompt + output + thoughts` when not + explicitly reported by the provider). + """ + + prompt_tokens: int = 0 + output_tokens: int = 0 + thoughts_tokens: int = 0 + cached_tokens: int = 0 + total_tokens: int = 0 + + @classmethod + def from_metadata( + cls, meta: types.GenerateContentResponseUsageMetadata | None + ) -> AdvisorUsage: + """Builds an `AdvisorUsage` snapshot from GenAI usage metadata. + + Args: + meta: Usage metadata from `LlmResponse.usage_metadata`, or `None`. + + Returns: + An `AdvisorUsage` populated from `meta`, or all-zero counts if `None`. + """ + if meta is None: + return cls() + buckets = TokenUsage.from_usage_metadata(meta) + prompt = max(0, buckets.input_tokens or 0) + output = max(0, buckets.candidate_output_tokens or 0) + thoughts = max(0, buckets.reasoning_output_tokens or 0) + cached = max(0, buckets.cache_read_input_tokens or 0) + raw_total = max(0, meta.total_token_count or 0) + total = raw_total or (prompt + output + thoughts) + return cls( + prompt_tokens=prompt, + output_tokens=output, + thoughts_tokens=thoughts, + cached_tokens=cached, + total_tokens=total, + ) + + def __add__(self, other: AdvisorUsage) -> AdvisorUsage: + if not isinstance(other, AdvisorUsage): + return NotImplemented + return AdvisorUsage( + prompt_tokens=self.prompt_tokens + other.prompt_tokens, + output_tokens=self.output_tokens + other.output_tokens, + thoughts_tokens=self.thoughts_tokens + other.thoughts_tokens, + cached_tokens=self.cached_tokens + other.cached_tokens, + total_tokens=self.total_tokens + other.total_tokens, + ) + + def to_dict(self) -> dict[str, int]: + """Returns token counts as a JSON-serializable dictionary.""" + return { + 'prompt_tokens': self.prompt_tokens, + 'output_tokens': self.output_tokens, + 'thoughts_tokens': self.thoughts_tokens, + 'cached_tokens': self.cached_tokens, + 'total_tokens': self.total_tokens, + } + + +@dataclass(frozen=True, kw_only=True) +class AdvisorResult: + """Outcome of a single advisor model consultation. + + Attributes: + text: Visible guidance text produced by the advisor model. + model: Configured model identifier on the advisor `BaseLlm`. + model_version: Provider-reported model version string, if available. + usage: Token usage snapshot for the consultation. + latency_ms: End-to-end wall-clock duration of the consultation in ms. + """ + + text: str + model: str + model_version: str | None + usage: AdvisorUsage + latency_ms: float + + +def resolve_thinking_level( + level: str | types.ThinkingLevel | None, +) -> types.ThinkingLevel | None: + """Maps a user-supplied thinking level to `types.ThinkingLevel`. + + Args: + level: One of `'minimal'`, `'low'`, `'medium'`, `'high'` + (case-insensitive), `'none'` / `'off'` / `''` / `None` to leave thinking + unset, or a `types.ThinkingLevel` enum value. + + Returns: + The corresponding `types.ThinkingLevel`, or `None` if disabled. + + Raises: + ValueError: If `level` is not a recognized thinking level. + """ + if level is None: + return None + if isinstance(level, types.ThinkingLevel): + if level == types.ThinkingLevel.THINKING_LEVEL_UNSPECIFIED: + return None + return level + if not isinstance(level, str): + raise ValueError( + f'Invalid advisor thinking_level {level!r}; expected a string or ' + 'types.ThinkingLevel.' + ) + key = level.strip().lower() + if key in ('', 'none', 'off'): + return None + if key not in _THINKING_LEVEL_MAP: + valid = sorted([*_THINKING_LEVEL_MAP.keys(), 'off']) + raise ValueError( + f'Invalid advisor thinking_level {level!r}; expected one of {valid}.' + ) + return _THINKING_LEVEL_MAP[key] + + +def resolve_advisor_llm(model: str | BaseLlm) -> BaseLlm: + """Resolves a model name or `BaseLlm` instance into a `BaseLlm`. + + Args: + model: Either an already-constructed `BaseLlm` or a model identifier + accepted by `LLMRegistry.new_llm` (for example `'gemini-2.5-pro'`). + + Returns: + A `BaseLlm` instance for the advisor model. + + Raises: + ValueError: If `model` is neither a `BaseLlm` nor a non-empty string. + """ + if isinstance(model, BaseLlm): + return model + if isinstance(model, str) and model.strip(): + return LLMRegistry.new_llm(model.strip()) + raise ValueError( + f'Invalid advisor_model {model!r}; expected a non-empty model string or ' + 'a BaseLlm instance.' + ) + + +def _build_request( + *, + llm: BaseLlm, + contents: Sequence[types.Content], + system_instruction: str, + thinking_level: types.ThinkingLevel | None, + max_output_tokens: int | None, + base_config: types.GenerateContentConfig | None, + clear_thinking_config: bool = False, +) -> LlmRequest: + """Constructs a tool-free `LlmRequest` for the advisor call.""" + config = ( + copy.deepcopy(base_config) + if base_config is not None + else types.GenerateContentConfig() + ) + config.system_instruction = system_instruction + config.tools = [] + config.tool_config = None + if max_output_tokens is not None: + config.max_output_tokens = max_output_tokens + if clear_thinking_config: + config.thinking_config = None + elif thinking_level is not None: + existing = config.thinking_config + config.thinking_config = types.ThinkingConfig( + thinking_level=thinking_level, + include_thoughts=existing.include_thoughts if existing else None, + ) + return LlmRequest( + model=llm.model, + contents=list(contents), + config=config, + ) + + +async def _collect( + llm: BaseLlm, + response_gen: AsyncGenerator[LlmResponse, None], + responses: list[LlmResponse], +) -> tuple[ + str, + str | None, + types.FinishReason | None, + AdvisorUsage, +]: + """Iterates `response_gen` and extracts visible text and metadata.""" + text_chunks: list[str] = [] + model_version: str | None = None + finish_reason: types.FinishReason | None = None + last_usage_meta: types.GenerateContentResponseUsageMetadata | None = None + + try: + async for response in response_gen: + responses.append(response) + if response.model_version: + model_version = response.model_version + if response.finish_reason: + finish_reason = response.finish_reason + elif _hit_output_cap(response.error_code): + finish_reason = types.FinishReason.MAX_TOKENS + if response.usage_metadata is not None: + # Take the last reading rather than summing across yields: adapters + # report cumulative token counts on the final non-partial response. + last_usage_meta = response.usage_metadata + if ( + response.error_code + and not _hit_output_cap(response.error_code) + and not _hit_output_cap(response.finish_reason) + ): + raise AdvisorError( + f'Advisor ({llm.model}) returned error {response.error_code}: ' + f'{response.error_message or "no message"}' + ) + if response.partial: + continue + if response.content and response.content.parts: + for part in response.content.parts: + if getattr(part, 'thought', False): + continue + if part.text: + text_chunks.append(part.text) + finally: + await response_gen.aclose() + + text = ''.join(text_chunks).strip() + usage = AdvisorUsage.from_metadata(last_usage_meta) + return text, model_version, finish_reason, usage + + +def _record_telemetry( + *, + agent_name: str, + elapsed_s: float, + request: LlmRequest, + responses: Sequence[LlmResponse], + error: Exception | None = None, +) -> None: + """Emits standard ADK OpenTelemetry client duration and token metrics.""" + try: + # pylint: disable=protected-access + if ( + tracing._instrumented_with_opentelemetry_instrumentation_google_genai() + and is_gemini_model(request.model) + ): + return + + normalized_responses: list[LlmResponse] = [] + last_usage_meta: types.GenerateContentResponseUsageMetadata | None = None + last_model_version: str | None = None + for resp in responses: + if resp.model_version: + last_model_version = resp.model_version + if resp.usage_metadata is not None: + last_usage_meta = resp.usage_metadata + if responses: + tail = responses[-1].model_copy( + update={ + 'model_version': ( + last_model_version or responses[-1].model_version + ), + 'usage_metadata': ( + last_usage_meta + if last_usage_meta is not None + else responses[-1].usage_metadata + ), + } + ) + normalized_responses = [*responses[:-1], tail] + + _metrics.record_client_operation_duration( + agent_name=agent_name, + elapsed_s=elapsed_s, + llm_request=request, + responses=normalized_responses, + error=error, + ) + if last_usage_meta is not None and normalized_responses: + _metrics.record_client_token_usage( + agent_name=agent_name, + llm_request=request, + responses=normalized_responses, + ) + except Exception: # pylint: disable=broad-exception-caught + logger.debug( + 'Failed to record telemetry for advisor call (%s).', + request.model, + exc_info=True, + ) + + +async def call_advisor( + llm: BaseLlm, + contents: Sequence[types.Content], + *, + system_instruction: str, + thinking_level: types.ThinkingLevel | None = None, + max_output_tokens: int | None = None, + timeout_seconds: float | None = None, + generate_content_config: types.GenerateContentConfig | None = None, + agent_name: str = 'model_consult', +) -> AdvisorResult: + """Executes one non-streaming advisor call and returns its guidance text. + + If `thinking_level` (or `generate_content_config.thinking_config`) is set + and the underlying model rejects `thinking_config` (for example a model or + third-party adapter that does not support thinking levels), the call retries + once without `thinking_config`. + + Args: + llm: The resolved advisor `BaseLlm` instance. + contents: Handover conversation contents from `build_advisor_contents`. + system_instruction: System prompt instructing the advisor how to respond. + thinking_level: Optional `types.ThinkingLevel` for the advisor call. + max_output_tokens: Optional cap on total generated tokens (thoughts + text). + timeout_seconds: Optional positive wall-clock timeout in seconds. + generate_content_config: Optional base `GenerateContentConfig` to copy. + agent_name: Agent attribute recorded on OpenTelemetry client metrics. + + Returns: + An `AdvisorResult` with the advisor's visible response text, model info, + token usage, and wall-clock latency in milliseconds. + + Raises: + ValueError: If `timeout_seconds` is less than or equal to zero. + AdvisorError: If the call times out, errors, or produces no visible text. + """ + if timeout_seconds is not None and timeout_seconds <= 0: + raise ValueError( + f'timeout_seconds must be positive; got {timeout_seconds!r}.' + ) + + req = _build_request( + llm=llm, + contents=contents, + system_instruction=system_instruction, + thinking_level=thinking_level, + max_output_tokens=max_output_tokens, + base_config=generate_content_config, + ) + can_retry_without_thinking = req.config.thinking_config is not None + call_t0 = time.perf_counter() + deadline = call_t0 + timeout_seconds if timeout_seconds is not None else None + + while True: + attempt_t0 = time.perf_counter() + responses: list[LlmResponse] = [] + try: + coro = _collect( + llm, + llm.generate_content_async(req, stream=False), + responses, + ) + if deadline is not None: + remaining_timeout = max(0.0, deadline - time.perf_counter()) + text, model_version, finish_reason, usage = await asyncio.wait_for( + coro, timeout=remaining_timeout + ) + else: + text, model_version, finish_reason, usage = await coro + + effective_max_tokens = req.config.max_output_tokens + if not text and _hit_output_cap(finish_reason): + raise AdvisorError( + f'Advisor ({llm.model}) produced no visible text before hitting ' + f'max_output_tokens={effective_max_tokens} (thoughts consumed ' + f'{usage.thoughts_tokens} tokens). Increase max_output_tokens or ' + 'lower thinking_level.' + ) + + if not text: + raise AdvisorError( + f'Advisor ({llm.model}) returned an empty response ' + f'(finish_reason={finish_reason}).' + ) + break + # Before Python 3.11, asyncio.TimeoutError is not the builtin + # TimeoutError, so catch both to cover asyncio and transport timeouts. + except (asyncio.TimeoutError, TimeoutError) as exc: + _record_telemetry( + agent_name=agent_name, + elapsed_s=time.perf_counter() - attempt_t0, + request=req, + responses=responses, + error=exc, + ) + if timeout_seconds is not None: + raise AdvisorError( + f'Advisor ({llm.model}) timed out after {timeout_seconds}s.' + ) from exc + raise AdvisorError(f'Advisor ({llm.model}) timed out: {exc}') from exc + except AdvisorError as exc: + _record_telemetry( + agent_name=agent_name, + elapsed_s=time.perf_counter() - attempt_t0, + request=req, + responses=responses, + error=exc, + ) + raise + except Exception as exc: # pylint: disable=broad-exception-caught + _record_telemetry( + agent_name=agent_name, + elapsed_s=time.perf_counter() - attempt_t0, + request=req, + responses=responses, + error=exc, + ) + if can_retry_without_thinking and _is_thinking_config_error(exc): + can_retry_without_thinking = False + logger.info( + 'Advisor model %s rejected thinking_config (%s); retrying without ' + 'thinking_config.', + llm.model, + exc, + ) + req = _build_request( + llm=llm, + contents=contents, + system_instruction=system_instruction, + thinking_level=None, + max_output_tokens=max_output_tokens, + base_config=generate_content_config, + clear_thinking_config=True, + ) + continue + raise AdvisorError(f'Advisor ({llm.model}) call failed: {exc}') from exc + + _record_telemetry( + agent_name=agent_name, + elapsed_s=time.perf_counter() - attempt_t0, + request=req, + responses=responses, + ) + latency_ms = (time.perf_counter() - call_t0) * 1000.0 + + if _hit_output_cap(finish_reason): + text = f'{text}\n\n[advisor guidance truncated at max_output_tokens]' + + return AdvisorResult( + text=text, + model=llm.model, + model_version=model_version, + usage=usage, + latency_ms=latency_ms, + ) + + +def _hit_output_cap( + finish_reason: types.FinishReason | str | None, +) -> bool: + """Returns True if generation stopped because `MAX_TOKENS` was reached.""" + if finish_reason is None: + return False + return str(finish_reason).upper().endswith('MAX_TOKENS') + + +def _is_thinking_config_error(exc: BaseException) -> bool: + """Heuristic check for errors caused by an unsupported `thinking_config`.""" + msg = str(exc).lower() + if any( + field in msg + for field in ( + 'thinking_config', + 'thinkingconfig', + 'thinking_level', + 'thinking level', + 'thinking_budget', + 'thinking budget', + ) + ): + return True + return 'thinking' in msg and any( + token in msg + for token in ( + 'unsupported', + 'not support', + 'unknown', + 'unexpected', + 'only available', + 'not allowed', + 'cannot', + ) + ) diff --git a/tests/unittests/tools/model_consult/test_advisor.py b/tests/unittests/tools/model_consult/test_advisor.py new file mode 100644 index 00000000000..1b414ff847e --- /dev/null +++ b/tests/unittests/tools/model_consult/test_advisor.py @@ -0,0 +1,730 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Unit tests for `google.adk.tools.model_consult._advisor`.""" + +from __future__ import annotations + +import asyncio +from typing import AsyncGenerator +from unittest import mock + +from google.adk.models.base_llm import BaseLlm +from google.adk.models.llm_request import LlmRequest +from google.adk.models.llm_response import LlmResponse +from google.adk.models.registry import LLMRegistry +from google.adk.telemetry import _metrics +from google.adk.telemetry import tracing +from google.adk.tools.model_consult._advisor import AdvisorError +from google.adk.tools.model_consult._advisor import AdvisorUsage +from google.adk.tools.model_consult._advisor import call_advisor +from google.adk.tools.model_consult._advisor import resolve_advisor_llm +from google.adk.tools.model_consult._advisor import resolve_thinking_level +from google.genai import types +from pydantic import Field +import pytest + + +class _FakeAdvisorLlm(BaseLlm): + """Test double for `BaseLlm` yielding scripted responses.""" + + model: str = 'fake-advisor-pro' + scripted_outcomes: list[list[LlmResponse] | BaseException] = Field( + default_factory=list + ) + recorded_requests: list[LlmRequest] = Field(default_factory=list) + recorded_streams: list[bool] = Field(default_factory=list) + delay_seconds: float = 0.0 + + async def generate_content_async( + self, llm_request: LlmRequest, stream: bool = False + ) -> AsyncGenerator[LlmResponse, None]: + self.recorded_requests.append(llm_request.model_copy(deep=True)) + self.recorded_streams.append(stream) + if self.delay_seconds > 0: + await asyncio.sleep(self.delay_seconds) + if not self.scripted_outcomes: + return + outcome = self.scripted_outcomes.pop(0) + if isinstance(outcome, BaseException): + raise outcome + for resp in outcome: + yield resp + + +def _sample_contents() -> tuple[types.Content, ...]: + return ( + types.Content( + role='user', + parts=[ + types.Part.from_text(text='How should I structure this retry?') + ], + ), + ) + + +@pytest.mark.parametrize( + ('raw_level', 'expected'), + [ + (None, None), + ('', None), + (' ', None), + ('none', None), + ('OFF', None), + ('minimal', types.ThinkingLevel.MINIMAL), + ('LOW', types.ThinkingLevel.LOW), + (' Medium ', types.ThinkingLevel.MEDIUM), + ('high', types.ThinkingLevel.HIGH), + (types.ThinkingLevel.HIGH, types.ThinkingLevel.HIGH), + (types.ThinkingLevel.THINKING_LEVEL_UNSPECIFIED, None), + ], +) +def test_resolve_thinking_level_valid( + raw_level: str | types.ThinkingLevel | None, + expected: types.ThinkingLevel | None, +): + """Normalizes valid thinking level strings, enums, and off/none values.""" + assert resolve_thinking_level(raw_level) == expected + + +@pytest.mark.parametrize('bad_level', ['ultra', 'maximum', 42]) +def test_resolve_thinking_level_invalid_raises(bad_level): + """Raises ValueError when given an unrecognized thinking level.""" + with pytest.raises(ValueError, match='Invalid advisor thinking_level'): + resolve_thinking_level(bad_level) + + +def test_resolve_advisor_llm_passes_through_instance(): + """Returns an already-constructed BaseLlm instance unchanged.""" + llm = _FakeAdvisorLlm() + assert resolve_advisor_llm(llm) is llm + + +def test_resolve_advisor_llm_resolves_string_via_registry(): + """Strips and resolves a model string via LLMRegistry.new_llm.""" + fake_llm = _FakeAdvisorLlm() + with mock.patch.object( + LLMRegistry, 'new_llm', autospec=True, return_value=fake_llm + ) as mock_new_llm: + resolved = resolve_advisor_llm(' gemini-2.5-pro ') + assert resolved is fake_llm + mock_new_llm.assert_called_once_with('gemini-2.5-pro') + + +@pytest.mark.parametrize('bad_model', ['', ' ', None]) +def test_resolve_advisor_llm_invalid_raises(bad_model): + """Raises ValueError when advisor_model is empty or not a string/BaseLlm.""" + with pytest.raises(ValueError, match='Invalid advisor_model'): + resolve_advisor_llm(bad_model) # type: ignore[arg-type] + + +def test_advisor_usage_from_metadata_and_addition(): + """Computes token totals, clamps negative sentinels, and adds snapshots.""" + assert AdvisorUsage.from_metadata(None) == AdvisorUsage() + + meta_fallback_total = types.GenerateContentResponseUsageMetadata( + prompt_token_count=100, + tool_use_prompt_token_count=15, + candidates_token_count=40, + thoughts_token_count=60, + cached_content_token_count=25, + total_token_count=None, + ) + u1 = AdvisorUsage.from_metadata(meta_fallback_total) + assert u1 == AdvisorUsage( + prompt_tokens=115, + output_tokens=40, + thoughts_tokens=60, + cached_tokens=25, + total_tokens=215, + ) + + meta_explicit_total = types.GenerateContentResponseUsageMetadata( + prompt_token_count=10, + candidates_token_count=5, + thoughts_token_count=2, + cached_content_token_count=-1, + total_token_count=50, + ) + u2 = AdvisorUsage.from_metadata(meta_explicit_total) + assert u2 == AdvisorUsage( + prompt_tokens=10, + output_tokens=5, + thoughts_tokens=2, + cached_tokens=0, + total_tokens=50, + ) + + combined = u1 + u2 + assert combined.to_dict() == { + 'prompt_tokens': 125, + 'output_tokens': 45, + 'thoughts_tokens': 62, + 'cached_tokens': 25, + 'total_tokens': 265, + } + with pytest.raises(TypeError): + _ = u1 + 'invalid' # type: ignore[operator] + + +@pytest.mark.asyncio +async def test_call_advisor_happy_path_filters_thoughts_and_partials(): + """Collects visible text across chunks and records OTel metrics.""" + base_cfg = types.GenerateContentConfig( + temperature=0.2, + max_output_tokens=1024, + tool_config=types.ToolConfig( + function_calling_config=types.FunctionCallingConfig( + mode=types.FunctionCallingConfigMode.ANY + ) + ), + thinking_config=types.ThinkingConfig(include_thoughts=True), + ) + llm = _FakeAdvisorLlm( + scripted_outcomes=[[ + LlmResponse( + partial=True, + content=types.Content( + role='model', + parts=[types.Part.from_text(text='partial duplicate')], + ), + usage_metadata=types.GenerateContentResponseUsageMetadata( + prompt_token_count=999, + candidates_token_count=5, + total_token_count=1004, + ), + ), + LlmResponse( + partial=False, + model_version='gemini-2.5-pro-001', + finish_reason=types.FinishReason.STOP, + content=types.Content( + role='model', + parts=[ + types.Part(text='internal thought', thought=True), + types.Part.from_text(text=' Use exponential '), + ], + ), + usage_metadata=types.GenerateContentResponseUsageMetadata( + prompt_token_count=50, + candidates_token_count=20, + thoughts_token_count=30, + total_token_count=100, + ), + ), + LlmResponse( + partial=False, + model_version=None, + finish_reason=None, + content=types.Content( + role='model', + parts=[types.Part.from_text(text='backoff. ')], + ), + usage_metadata=None, + ), + ]] + ) + + with ( + mock.patch.object( + _metrics, 'record_client_operation_duration', autospec=True + ) as mock_duration, + mock.patch.object( + _metrics, 'record_client_token_usage', autospec=True + ) as mock_tokens, + ): + result = await call_advisor( + llm, + _sample_contents(), + system_instruction='Give concise advice.', + thinking_level=types.ThinkingLevel.HIGH, + generate_content_config=base_cfg, + ) + + assert result.text == 'Use exponential backoff.' + assert result.model == 'fake-advisor-pro' + assert result.model_version == 'gemini-2.5-pro-001' + assert result.usage == AdvisorUsage( + prompt_tokens=50, + output_tokens=20, + thoughts_tokens=30, + cached_tokens=0, + total_tokens=100, + ) + assert result.latency_ms > 0.0 + + assert llm.recorded_streams == [False] + assert len(llm.recorded_requests) == 1 + sent_cfg = llm.recorded_requests[0].config + assert sent_cfg.system_instruction == 'Give concise advice.' + assert sent_cfg.tools == [] + assert sent_cfg.tool_config is None + assert sent_cfg.max_output_tokens == 1024 + assert sent_cfg.temperature == 0.2 + assert sent_cfg.thinking_config.thinking_level == types.ThinkingLevel.HIGH + assert sent_cfg.thinking_config.include_thoughts is True + assert base_cfg.thinking_config.thinking_level is None + + mock_duration.assert_called_once() + assert mock_duration.call_args.kwargs['agent_name'] == 'model_consult' + assert mock_duration.call_args.kwargs['error'] is None + assert ( + mock_duration.call_args.kwargs['responses'][-1].model_version + == 'gemini-2.5-pro-001' + ) + mock_tokens.assert_called_once() + assert mock_tokens.call_args.kwargs['agent_name'] == 'model_consult' + assert ( + mock_tokens.call_args.kwargs['responses'][ + -1 + ].usage_metadata.total_token_count + == 100 + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + 'err_msg', + [ + 'thinking_config is unsupported for this model', + 'thinking_level is not supported', + 'Model claude-3-5-haiku does not support thinking', + ( + 'thinking_budget must be set explicitly when ThinkingConfig is ' + 'provided for Anthropic models' + ), + 'Thinking is only available on Gemini 2.5 and newer models', + ], +) +async def test_call_advisor_retries_without_thinking_config_on_rejection( + err_msg: str, +): + """Retries once without thinking_config and records telemetry for both.""" + base_cfg = types.GenerateContentConfig( + thinking_config=types.ThinkingConfig( + thinking_level=types.ThinkingLevel.HIGH + ) + ) + llm = _FakeAdvisorLlm( + scripted_outcomes=[ + ValueError(err_msg), + [ + LlmResponse( + finish_reason=types.FinishReason.STOP, + content=types.Content( + role='model', + parts=[types.Part.from_text(text='Fallback succeeded.')], + ), + ) + ], + ] + ) + + with mock.patch.object( + _metrics, 'record_client_operation_duration', autospec=True + ) as mock_duration: + result = await call_advisor( + llm, + _sample_contents(), + system_instruction='Advisor system prompt.', + thinking_level=None, + generate_content_config=base_cfg, + ) + + assert result.text == 'Fallback succeeded.' + assert len(llm.recorded_requests) == 2 + assert llm.recorded_requests[0].config.thinking_config is not None + assert llm.recorded_requests[1].config.thinking_config is None + assert mock_duration.call_count == 2 + assert isinstance(mock_duration.call_args_list[0].kwargs['error'], ValueError) + assert mock_duration.call_args_list[1].kwargs['error'] is None + + +@pytest.mark.asyncio +async def test_call_advisor_preserves_or_overrides_caller_thinking_budget(): + """Preserves thinking_budget when thinking_level=None; overrides when set.""" + base_cfg = types.GenerateContentConfig( + thinking_config=types.ThinkingConfig( + thinking_budget=2048, include_thoughts=True + ) + ) + llm = _FakeAdvisorLlm( + scripted_outcomes=[ + [ + LlmResponse( + finish_reason=types.FinishReason.STOP, + content=types.Content( + role='model', + parts=[types.Part.from_text(text='Used budget.')], + ), + ) + ], + [ + LlmResponse( + finish_reason=types.FinishReason.STOP, + content=types.Content( + role='model', + parts=[types.Part.from_text(text='Used level.')], + ), + ) + ], + ] + ) + + result_preserved = await call_advisor( + llm, + _sample_contents(), + system_instruction='Advisor system prompt.', + thinking_level=None, + generate_content_config=base_cfg, + ) + assert result_preserved.text == 'Used budget.' + sent_preserved = llm.recorded_requests[0].config.thinking_config + assert sent_preserved.thinking_budget == 2048 + assert sent_preserved.thinking_level is None + + result_overridden = await call_advisor( + llm, + _sample_contents(), + system_instruction='Advisor system prompt.', + thinking_level=types.ThinkingLevel.HIGH, + generate_content_config=base_cfg, + ) + assert result_overridden.text == 'Used level.' + sent_overridden = llm.recorded_requests[1].config.thinking_config + assert sent_overridden.thinking_level == types.ThinkingLevel.HIGH + assert sent_overridden.thinking_budget is None + assert sent_overridden.include_thoughts is True + + +@pytest.mark.asyncio +async def test_call_advisor_does_not_retry_unrelated_invalid_argument_errors(): + """Does not retry 400 INVALID_ARGUMENT errors unrelated to thinking config.""" + llm = _FakeAdvisorLlm( + scripted_outcomes=[ + RuntimeError( + '400 INVALID_ARGUMENT: Invalid value at contents[0] ' + '(text: "I am thinking about this")' + ), + [ + LlmResponse( + content=types.Content( + role='model', + parts=[types.Part.from_text(text='Should not run')], + ) + ) + ], + ] + ) + + with pytest.raises(AdvisorError, match='400 INVALID_ARGUMENT'): + await call_advisor( + llm, + _sample_contents(), + system_instruction='Advisor system prompt.', + thinking_level=types.ThinkingLevel.HIGH, + ) + assert len(llm.recorded_requests) == 1 + + +@pytest.mark.asyncio +async def test_call_advisor_max_tokens_with_no_visible_text_raises(): + """Raises thought-starvation AdvisorError even with error_code=MAX_TOKENS.""" + base_cfg = types.GenerateContentConfig(max_output_tokens=512) + llm = _FakeAdvisorLlm( + scripted_outcomes=[[ + LlmResponse( + finish_reason=types.FinishReason.MAX_TOKENS, + error_code=types.FinishReason.MAX_TOKENS, + content=types.Content(role='model', parts=[]), + usage_metadata=types.GenerateContentResponseUsageMetadata( + prompt_token_count=100, + thoughts_token_count=512, + total_token_count=612, + ), + ) + ]] + ) + + with mock.patch.object( + _metrics, 'record_client_operation_duration', autospec=True + ) as mock_duration: + with pytest.raises( + AdvisorError, + match=( + r'no visible text before hitting max_output_tokens=512.*512 tokens' + ), + ): + await call_advisor( + llm, + _sample_contents(), + system_instruction='Advisor system prompt.', + generate_content_config=base_cfg, + ) + mock_duration.assert_called_once() + assert isinstance(mock_duration.call_args.kwargs['error'], AdvisorError) + + +@pytest.mark.asyncio +async def test_call_advisor_max_tokens_with_partial_text_and_error_code(): + """Returns truncated text when LiteLlm sets error_code=MAX_TOKENS.""" + llm = _FakeAdvisorLlm( + scripted_outcomes=[[ + LlmResponse( + finish_reason=types.FinishReason.MAX_TOKENS, + error_code=types.FinishReason.MAX_TOKENS, + error_message='Maximum tokens reached', + content=types.Content( + role='model', + parts=[types.Part.from_text(text='Step 1: check logs.')], + ), + ) + ]] + ) + + result = await call_advisor( + llm, + _sample_contents(), + system_instruction='Advisor system prompt.', + max_output_tokens=64, + ) + assert result.text == ( + 'Step 1: check logs.\n\n[advisor guidance truncated at max_output_tokens]' + ) + + +@pytest.mark.asyncio +async def test_call_advisor_empty_response_on_stop_raises(): + """Raises AdvisorError when finish_reason is STOP/None with empty text.""" + llm = _FakeAdvisorLlm( + scripted_outcomes=[[ + LlmResponse( + finish_reason=None, + content=types.Content( + role='model', + parts=[types.Part.from_text(text=' ')], + ), + ) + ]] + ) + + with pytest.raises(AdvisorError, match='returned an empty response'): + await call_advisor( + llm, + _sample_contents(), + system_instruction='Advisor system prompt.', + ) + + +@pytest.mark.asyncio +async def test_call_advisor_response_error_code_raises_and_records_telemetry(): + """Raises AdvisorError on error_code and preserves responses for telemetry.""" + llm = _FakeAdvisorLlm( + scripted_outcomes=[[ + LlmResponse( + model_version='gemini-2.5-pro-002', + error_code='RESOURCE_EXHAUSTED', + error_message=None, + ) + ]] + ) + + with mock.patch.object( + _metrics, 'record_client_operation_duration', autospec=True + ) as mock_duration: + with pytest.raises( + AdvisorError, match='returned error RESOURCE_EXHAUSTED: no message' + ): + await call_advisor( + llm, + _sample_contents(), + system_instruction='Advisor system prompt.', + ) + mock_duration.assert_called_once() + assert ( + mock_duration.call_args.kwargs['responses'][-1].model_version + == 'gemini-2.5-pro-002' + ) + + +@pytest.mark.asyncio +async def test_call_advisor_telemetry_failure_does_not_break_call(): + """Swallows telemetry recording errors so advisor calls still succeed.""" + llm = _FakeAdvisorLlm( + scripted_outcomes=[[ + LlmResponse( + finish_reason=types.FinishReason.STOP, + content=types.Content( + role='model', + parts=[types.Part.from_text(text='Still works.')], + ), + ) + ]] + ) + with mock.patch.object( + _metrics, + 'record_client_operation_duration', + autospec=True, + side_effect=RuntimeError('OTel exporter error'), + ): + result = await call_advisor( + llm, + _sample_contents(), + system_instruction='Advisor system prompt.', + ) + assert result.text == 'Still works.' + + +@pytest.mark.asyncio +async def test_call_advisor_timeout_raises_advisor_error(): + """Raises AdvisorError when the advisor call exceeds timeout_seconds.""" + llm = _FakeAdvisorLlm(delay_seconds=0.2) + with pytest.raises(AdvisorError, match='timed out after 0.01s'): + await call_advisor( + llm, + _sample_contents(), + system_instruction='Advisor system prompt.', + timeout_seconds=0.01, + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize('bad_timeout', [0, 0.0, -5.0]) +async def test_call_advisor_non_positive_timeout_raises_value_error( + bad_timeout: float, +): + """Rejects zero or negative timeout_seconds with ValueError.""" + llm = _FakeAdvisorLlm() + with pytest.raises(ValueError, match='timeout_seconds must be positive'): + await call_advisor( + llm, + _sample_contents(), + system_instruction='Advisor system prompt.', + timeout_seconds=bad_timeout, + ) + + +@pytest.mark.asyncio +async def test_call_advisor_transport_timeout_without_timeout_seconds(): + """Formats transport TimeoutError without 'Nones' when timeout is None.""" + llm = _FakeAdvisorLlm( + scripted_outcomes=[TimeoutError('read timed out on socket')] + ) + with pytest.raises( + AdvisorError, + match=r'Advisor \(fake-advisor-pro\) timed out: read timed out on socket', + ) as exc_info: + await call_advisor( + llm, + _sample_contents(), + system_instruction='Advisor system prompt.', + timeout_seconds=None, + ) + assert 'Nones' not in str(exc_info.value) + + +@pytest.mark.asyncio +async def test_call_advisor_timeout_bounds_total_wall_clock_across_retry(): + """Shares timeout_seconds budget across the initial attempt and retry.""" + llm = _FakeAdvisorLlm( + delay_seconds=0.04, + scripted_outcomes=[ + ValueError('thinking_config is unsupported for this model'), + [ + LlmResponse( + finish_reason=types.FinishReason.STOP, + content=types.Content( + role='model', + parts=[types.Part.from_text(text='Too slow.')], + ), + ) + ], + ], + ) + with pytest.raises(AdvisorError, match='timed out after 0.06s'): + await call_advisor( + llm, + _sample_contents(), + system_instruction='Advisor system prompt.', + thinking_level=types.ThinkingLevel.HIGH, + timeout_seconds=0.06, + ) + + +@pytest.mark.asyncio +async def test_call_advisor_skips_native_telemetry_when_genai_instrumented(): + """Skips native OTel metrics for Gemini when genai OTel lib is active.""" + gemini_llm = _FakeAdvisorLlm( + model='gemini-2.5-pro', + scripted_outcomes=[[ + LlmResponse( + finish_reason=types.FinishReason.STOP, + content=types.Content( + role='model', + parts=[types.Part.from_text(text='Gemini advice.')], + ), + usage_metadata=types.GenerateContentResponseUsageMetadata( + prompt_token_count=10, + candidates_token_count=5, + total_token_count=15, + ), + ) + ]], + ) + non_gemini_llm = _FakeAdvisorLlm( + model='claude-3-7-sonnet', + scripted_outcomes=[[ + LlmResponse( + finish_reason=types.FinishReason.STOP, + content=types.Content( + role='model', + parts=[types.Part.from_text(text='Claude advice.')], + ), + usage_metadata=types.GenerateContentResponseUsageMetadata( + prompt_token_count=10, + candidates_token_count=5, + total_token_count=15, + ), + ) + ]], + ) + + with ( + mock.patch.object( + tracing, + '_instrumented_with_opentelemetry_instrumentation_google_genai', + return_value=True, + ), + mock.patch.object( + _metrics, 'record_client_operation_duration', autospec=True + ) as mock_duration, + mock.patch.object( + _metrics, 'record_client_token_usage', autospec=True + ) as mock_tokens, + ): + await call_advisor( + gemini_llm, + _sample_contents(), + system_instruction='Advisor system prompt.', + ) + mock_duration.assert_not_called() + mock_tokens.assert_not_called() + + await call_advisor( + non_gemini_llm, + _sample_contents(), + system_instruction='Advisor system prompt.', + ) + mock_duration.assert_called_once() + mock_tokens.assert_called_once() From 84cc99abd53010817a85b51fdcdc5d4f48366769 Mon Sep 17 00:00:00 2001 From: Xuan Yang Date: Tue, 29 Sep 2026 17:00:05 -0700 Subject: [PATCH 25/29] feat(tools): add ModelConsultTool with turn and session budgets Add `ModelConsultTool`, an ADK tool that lets an executor `LlmAgent` consult a stronger advisor model mid-generation without relinquishing control of the conversation. Key capabilities: - Per-turn (`max_uses`) and session-wide (`session_max_uses`) consult budgets stored in session state (`temp:` and persistent state keys) with `has_remaining_budget` helper and per-session concurrency/delta coordination for parallel tool calls. - Automatic context handover via `build_advisor_contents` (with both `'events'` and `'transcript'` handover modes on `ModelConsultContextConfig`). - Forwards the executor's `static_instruction`, state-interpolated `canonical_instruction`, and non-self tool inventory (`canonical_tools`) to the advisor system instruction. - Graceful degradation on budget exhaustion (`status='limit_reached'`), missing question (`status='invalid_request'`), and advisor runtime or timeout errors (`status='error'`). - Top-level lazy export of `ModelConsultTool` from `google.adk.tools`. Co-authored-by: Xuan Yang PiperOrigin-RevId: 990614577 --- src/google/adk/tools/__init__.py | 10 + .../adk/tools/model_consult/__init__.py | 14 + .../adk/tools/model_consult/_context.py | 18 +- .../model_consult/_model_consult_tool.py | 818 ++++++++++ .../adk/tools/model_consult/_prompts.py | 142 ++ .../model_consult/test_model_consult_tool.py | 1354 +++++++++++++++++ 6 files changed, 2353 insertions(+), 3 deletions(-) create mode 100644 src/google/adk/tools/model_consult/_model_consult_tool.py create mode 100644 src/google/adk/tools/model_consult/_prompts.py create mode 100644 tests/unittests/tools/model_consult/test_model_consult_tool.py diff --git a/src/google/adk/tools/__init__.py b/src/google/adk/tools/__init__.py index 9cd22f305ff..b825084ef8f 100644 --- a/src/google/adk/tools/__init__.py +++ b/src/google/adk/tools/__init__.py @@ -38,6 +38,8 @@ from .load_artifacts_tool import load_artifacts_tool as load_artifacts from .load_memory_tool import load_memory_tool as load_memory from .long_running_tool import LongRunningFunctionTool + from .model_consult import ModelConsultContextConfig + from .model_consult import ModelConsultTool from .preload_memory_tool import preload_memory_tool as preload_memory from .tool_context import ToolContext from .transfer_to_agent_tool import transfer_to_agent @@ -82,6 +84,14 @@ '.long_running_tool', 'LongRunningFunctionTool', ), + 'ModelConsultContextConfig': ( + '.model_consult._context', + 'ModelConsultContextConfig', + ), + 'ModelConsultTool': ( + '.model_consult._model_consult_tool', + 'ModelConsultTool', + ), 'preload_memory': ('.preload_memory_tool', 'preload_memory_tool'), 'request_input': ('._request_input_tool', 'request_input'), 'RemoteMcpServer': ('._remote_mcp_server', 'RemoteMcpServer'), diff --git a/src/google/adk/tools/model_consult/__init__.py b/src/google/adk/tools/model_consult/__init__.py index a469d29c802..272f37623b9 100644 --- a/src/google/adk/tools/model_consult/__init__.py +++ b/src/google/adk/tools/model_consult/__init__.py @@ -14,8 +14,22 @@ """Lets a fast executor model consult a stronger advisor model mid-task.""" +from ._context import ContextMode from ._context import ModelConsultContextConfig +from ._model_consult_tool import DEFAULT_ADVISOR_MODEL +from ._model_consult_tool import DEFAULT_TOOL_NAME +from ._model_consult_tool import ModelConsultTool +from ._prompts import ADVISOR_SYSTEM_INSTRUCTION +from ._prompts import EXECUTOR_INSTRUCTION +from ._prompts import TOOL_DESCRIPTION __all__ = [ + 'ADVISOR_SYSTEM_INSTRUCTION', + 'ContextMode', + 'DEFAULT_ADVISOR_MODEL', + 'DEFAULT_TOOL_NAME', + 'EXECUTOR_INSTRUCTION', 'ModelConsultContextConfig', + 'ModelConsultTool', + 'TOOL_DESCRIPTION', ] diff --git a/src/google/adk/tools/model_consult/_context.py b/src/google/adk/tools/model_consult/_context.py index 75f7a58496c..c110161a8b1 100644 --- a/src/google/adk/tools/model_consult/_context.py +++ b/src/google/adk/tools/model_consult/_context.py @@ -26,6 +26,7 @@ from collections.abc import Sequence import json from typing import Any +from typing import Literal from typing import TYPE_CHECKING from google.genai import types @@ -38,6 +39,8 @@ if TYPE_CHECKING: from ...events.event import Event +ContextMode = Literal['events', 'transcript'] + _ROLE_LABELS = {'user': 'USER', 'model': 'AGENT'} # How much more room plain text gets than a rendered tool payload. Documented # on ModelConsultContextConfig.max_part_chars. @@ -50,7 +53,15 @@ class ModelConsultContextConfig(BaseModel): """Controls how much of the executor's session reaches the advisor.""" - model_config = ConfigDict(extra='forbid') + model_config = ConfigDict(extra='forbid', use_attribute_docstrings=True) + + mode: ContextMode = 'events' + """How the session is shaped for the advisor. + + `'events'` hands over multi-turn `types.Content` objects; `'transcript'` + collapses the session into one labelled plain-text block inside the final + user message. + """ include_session: bool = True """Whether to send the session at all. @@ -196,9 +207,10 @@ def _convert_part( return types.Part(text=text) if part.inline_data is not None or part.file_data is not None: - if config.include_media: + if config.include_media and config.mode != 'transcript': return part - description = _describe_media_part(part, reason='omitted') + reason = '' if config.include_media else 'omitted' + description = _describe_media_part(part, reason=reason) return None if description is None else types.Part(text=description) if part.executable_code is not None: diff --git a/src/google/adk/tools/model_consult/_model_consult_tool.py b/src/google/adk/tools/model_consult/_model_consult_tool.py new file mode 100644 index 00000000000..e6af4242818 --- /dev/null +++ b/src/google/adk/tools/model_consult/_model_consult_tool.py @@ -0,0 +1,818 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""ModelConsultTool: mid-generation escalation from executor to advisor.""" + +from __future__ import annotations + +import asyncio +from collections.abc import Sequence +import dataclasses +import logging +from typing import Any +from typing import TYPE_CHECKING +import weakref + +from google.genai import types +from typing_extensions import override + +from ...agents.readonly_context import ReadonlyContext +from ...utils.instructions_utils import inject_session_state +from ..base_tool import BaseTool +from ._advisor import AdvisorError +from ._advisor import AdvisorResult +from ._advisor import call_advisor +from ._advisor import resolve_advisor_llm +from ._advisor import resolve_thinking_level +from ._context import _merge_adjacent +from ._context import build_advisor_contents +from ._context import ModelConsultContextConfig +from ._context import render_transcript +from ._prompts import ADVISOR_HANDOFF_TEMPLATE +from ._prompts import ADVISOR_SYSTEM_INSTRUCTION +from ._prompts import CONTEXT_BLOCK_TEMPLATE +from ._prompts import EXECUTOR_INSTRUCTION +from ._prompts import TOOL_DESCRIPTION + +if TYPE_CHECKING: + from ...agents.callback_context import CallbackContext + from ...events.event import Event + from ...models.base_llm import BaseLlm + from ...models.llm_request import LlmRequest + from ..tool_context import ToolContext + +logger = logging.getLogger('google_adk.' + __name__) + +DEFAULT_ADVISOR_MODEL = 'gemini-3.1-pro-preview' +DEFAULT_TOOL_NAME = 'model_consult' + +# `temp:` state is applied to the live session for the duration of an +# invocation and never persisted, and the invocation id in the key ensures the +# turn budget resets on the next turn even if a caller reuses a session object. +_TURN_USES_STATE_KEY_TEMPLATE = 'temp:model_consult:{name}:{invocation_id}:uses' +# Non-`temp:` state persists across turns in the same session so a multi-turn +# conversation cannot exceed `session_max_uses`. +_SESSION_USES_STATE_KEY_TEMPLATE = 'model_consult:{name}:session_uses' + +_TOOL_DESCRIPTION_LIMIT = 300 + +_TURN_LIMIT_MESSAGE = ( + 'The advisor consult budget for this turn is exhausted ({max_uses} of' + ' {max_uses} used). Continue with your own best judgment, reusing the' + ' guidance you already received.' +) + +_SESSION_LIMIT_MESSAGE = ( + 'The advisor consult budget for this session is exhausted' + ' ({session_max_uses} of {session_max_uses} used). Continue with your own' + ' best judgment, reusing the guidance you already received.' +) + +_ERROR_MESSAGE = ( + 'The advisor could not be reached. Continue with your own best judgment;' + ' do not retry this tool for the same question.' +) + + +@dataclasses.dataclass +class _SessionConsultState: + """Per-session concurrency and state-delta coordination for parallel calls.""" + + cond: asyncio.Condition = dataclasses.field(default_factory=asyncio.Condition) + session_ref: Any = None + active_calls: int = 0 + reserved_session: int = 0 + reserved_turn: dict[str, int] = dataclasses.field(default_factory=dict) + inv_deltas: dict[str, list[dict[str, Any]]] = dataclasses.field( + default_factory=dict + ) + + +def _extract_static_instruction_texts(value: Any) -> list[str]: + """Extracts text segments from a `types.ContentUnion` static instruction.""" + if isinstance(value, str): + return [value] + if isinstance(value, types.Part): + return [value.text] if value.text else [] + if isinstance(value, types.Content): + return [ + text + for part in value.parts or [] + for text in _extract_static_instruction_texts(part) + ] + if isinstance(value, Sequence) and not isinstance( + value, (str, bytes, bytearray) + ): + return [ + text + for item in value + for text in _extract_static_instruction_texts(item) + ] + return [] + + +class ModelConsultTool(BaseTool): + """Lets an executor agent consult a stronger advisor model mid-generation. + + The advisor reads the executor's full session -- instructions, reasoning, + tool calls and tool results -- and returns a plan or course correction. The + executor keeps doing the work, so the bulk of token generation stays at + executor rates. + + Example: + ```python + root_agent = Agent( + model='gemini-2.5-flash', + name='root_cause_analysis_agent', + instruction='...', + tools=[ + ModelConsultTool( + model='gemini-3.1-pro-preview', + max_uses=2, + session_max_uses=5, + ) + ], + ) + ``` + + Attributes: + advisor_model: The resolved advisor `BaseLlm`. + max_uses: Consults allowed per turn, or `None` for unlimited. + session_max_uses: Consults allowed across the entire session, or `None` for + unlimited. + max_output_tokens: Output token cap for the advisor response, if configured. + thinking_level: Normalized advisor thinking level (`minimal`, `low`, + `medium`, `high`), or `None`. + """ + + def __init__( + self, + *, + model: str | BaseLlm = DEFAULT_ADVISOR_MODEL, + max_uses: int | None = None, + session_max_uses: int | None = None, + thinking_level: str | types.ThinkingLevel | None = 'high', + max_output_tokens: int | None = None, + name: str = DEFAULT_TOOL_NAME, + description: str | None = None, + executor_instruction: str | None = None, + advisor_instruction: str | None = None, + include_agent_instruction: bool = True, + include_tool_inventory: bool = True, + context_config: ModelConsultContextConfig | None = None, + generate_content_config: types.GenerateContentConfig | None = None, + timeout_seconds: float | None = None, + ): + """Initializes the tool. + + Args: + model: Advisor model name (resolved through ADK's model registry) or a + `BaseLlm` instance, so any model ADK supports can advise. + max_uses: Maximum advisor consults allowed in a single turn. `None` (the + default) means unlimited. + session_max_uses: Maximum advisor consults allowed across the entire + session. `None` (the default) means unlimited. When both `max_uses` and + `session_max_uses` are set, whichever limit is reached first blocks + further consults. + thinking_level: `'minimal'`, `'low'`, `'medium'`, `'high'` (default), or + a `types.ThinkingLevel` enum value. `None` disables overriding the + thinking config. + max_output_tokens: Caps the advisor's output (thinking included on models + that bill it there). Advisor output is the single largest cost driver of + this pattern, so a cap is the cheapest lever available -- but a cap that + is too tight starves a high thinking_level and returns nothing. Measure + before setting it below ~4096 with `thinking_level='high'`. + name: Tool name the executor sees. Change it only if it collides. + description: Overrides the tuned tool description that steers escalation. + executor_instruction: Overrides the escalation policy automatically + appended to the executor's `system_instruction`. Pass `""` to disable + automatic injection. + advisor_instruction: Overrides the advisor's system instruction. + include_agent_instruction: Forward the executor agent's own instruction to + the advisor, so guidance respects the executor's constraints. + include_tool_inventory: Tell the advisor which tools the executor can + call, so the plan names real tools with real arguments instead of steps + the executor cannot perform. + context_config: How much session context to hand over. + generate_content_config: Extra generation config for the advisor call + (temperature, max_output_tokens, safety settings...). + timeout_seconds: Abort the advisor call after this long. On timeout the + executor is told to proceed on its own rather than failing the turn. + + Raises: + ValueError: If `max_uses`, `session_max_uses`, `max_output_tokens`, or + `timeout_seconds` is not positive, or if `thinking_level` or `model` is + invalid. + """ + super().__init__( + name=name, + description=description or TOOL_DESCRIPTION, + ) + if max_uses is not None and max_uses <= 0: + raise ValueError(f'max_uses must be positive or None, got {max_uses}') + if session_max_uses is not None and session_max_uses <= 0: + raise ValueError( + f'session_max_uses must be positive or None, got {session_max_uses}' + ) + cfg_max_output_tokens = ( + generate_content_config.max_output_tokens + if generate_content_config is not None + else None + ) + if max_output_tokens is not None and max_output_tokens <= 0: + raise ValueError( + f'max_output_tokens must be positive or None, got {max_output_tokens}' + ) + if cfg_max_output_tokens is not None and cfg_max_output_tokens <= 0: + raise ValueError( + 'generate_content_config.max_output_tokens must be positive or None,' + f' got {cfg_max_output_tokens}' + ) + if ( + max_output_tokens is not None + and cfg_max_output_tokens is not None + and max_output_tokens != cfg_max_output_tokens + ): + raise ValueError( + f'Conflicting max_output_tokens ({max_output_tokens}) and' + ' generate_content_config.max_output_tokens' + f' ({cfg_max_output_tokens})' + ) + if timeout_seconds is not None and timeout_seconds <= 0: + raise ValueError( + f'timeout_seconds must be positive or None, got {timeout_seconds}' + ) + + self.advisor_model = resolve_advisor_llm(model) + self.max_output_tokens = ( + max_output_tokens + if max_output_tokens is not None + else cfg_max_output_tokens + ) + self.max_uses = max_uses + self.session_max_uses = session_max_uses + self._thinking_level = resolve_thinking_level(thinking_level) + self.thinking_level: str | None = ( + self._thinking_level.value.lower() + if self._thinking_level is not None + else None + ) + self._advisor_instruction = ( + advisor_instruction or ADVISOR_SYSTEM_INSTRUCTION + ) + self._include_agent_instruction = include_agent_instruction + self._include_tool_inventory = include_tool_inventory + self._context_config = context_config or ModelConsultContextConfig() + if max_output_tokens is not None and cfg_max_output_tokens is None: + generate_content_config = ( + generate_content_config.model_copy(deep=True) + if generate_content_config is not None + else types.GenerateContentConfig() + ) + generate_content_config.max_output_tokens = max_output_tokens + self._generate_content_config = generate_content_config + self._timeout_seconds = timeout_seconds + if executor_instruction is not None: + self._executor_system_instruction = executor_instruction.strip() + elif self.name == DEFAULT_TOOL_NAME: + self._executor_system_instruction = EXECUTOR_INSTRUCTION + else: + self._executor_system_instruction = EXECUTOR_INSTRUCTION.replace( + f'`{DEFAULT_TOOL_NAME}`', f'`{self.name}`' + ) + self._session_states: dict[int, _SessionConsultState] = {} + + def _get_declaration(self) -> types.FunctionDeclaration: + return types.FunctionDeclaration( + name=self.name, + description=self.description, + parameters=types.Schema( + type=types.Type.OBJECT, + properties={ + 'question': types.Schema( + type=types.Type.STRING, + description=( + 'The specific decision or blocker you want reviewed.' + ' State the approach you are considering, or what you' + ' tried and how it failed. Be concrete; the advisor' + ' already sees the conversation, so do not restate it.' + ), + ), + 'context': types.Schema( + type=types.Type.STRING, + description=( + 'Optional. Anything material that is NOT visible in the' + ' conversation: constraints you inferred, observations' + ' from outside this session, or the options you are' + ' weighing.' + ), + ), + }, + required=['question'], + ), + ) + + @override + async def process_llm_request( + self, + *, + tool_context: ToolContext, + llm_request: LlmRequest, + ) -> None: + await super().process_llm_request( + tool_context=tool_context, llm_request=llm_request + ) + if not self._executor_system_instruction: + return + existing = llm_request.config.system_instruction + if ( + not isinstance(existing, str) + or self._executor_system_instruction not in existing + ): + llm_request.append_instructions([self._executor_system_instruction]) + + def _turn_uses_state_key( + self, tool_context: ToolContext | CallbackContext + ) -> str: + return _TURN_USES_STATE_KEY_TEMPLATE.format( + name=self.name, + invocation_id=tool_context.invocation_id or 'unknown', + ) + + def _session_uses_state_key(self) -> str: + return _SESSION_USES_STATE_KEY_TEMPLATE.format(name=self.name) + + def _read_state_counter( + self, tool_context: ToolContext | CallbackContext, key: str + ) -> int: + try: + return max(int(tool_context.state.get(key, 0) or 0), 0) + except Exception: # pylint: disable=broad-exception-caught + logger.debug( + 'ModelConsultTool could not read state counter %s', + key, + exc_info=True, + ) + return 0 + + def _turn_uses_so_far( + self, tool_context: ToolContext | CallbackContext + ) -> int: + return self._read_state_counter( + tool_context, self._turn_uses_state_key(tool_context) + ) + + def _session_uses_so_far( + self, tool_context: ToolContext | CallbackContext + ) -> int: + return self._read_state_counter( + tool_context, self._session_uses_state_key() + ) + + def has_remaining_budget( + self, context: ToolContext | CallbackContext + ) -> bool: + """Returns whether at least one consult remains in the turn and session.""" + if ( + self.session_max_uses is not None + and self._session_uses_so_far(context) >= self.session_max_uses + ): + return False + if ( + self.max_uses is not None + and self._turn_uses_so_far(context) >= self.max_uses + ): + return False + return True + + def _write_state_counter( + self, tool_context: ToolContext, key: str, value: int + ) -> None: + try: + tool_context.state[key] = value + except TypeError: + # Fallback to the underlying `session.state` dict (`State._value`) if + # `tool_context.state` rejects item assignment (e.g., a custom state + # mapping or a `state_schema` validator that only exempts `app:`/`user:`/ + # `temp:` prefixes), while `_record_use` updates `actions.state_delta`. + tool_context.session.state[key] = value + + def _record_use( + self, + tool_context: ToolContext, + *, + turn_uses: int, + session_uses: int, + deltas: list[dict[str, Any]], + ) -> None: + turn_key = self._turn_uses_state_key(tool_context) + session_key = self._session_uses_state_key() + try: + self._write_state_counter(tool_context, turn_key, turn_uses) + self._write_state_counter(tool_context, session_key, session_uses) + for delta in deltas: + delta[turn_key] = max(int(delta.get(turn_key, 0) or 0), turn_uses) + delta[session_key] = max( + int(delta.get(session_key, 0) or 0), session_uses + ) + except Exception: # pylint: disable=broad-exception-caught + logger.warning( + 'ModelConsultTool could not persist its use counters', exc_info=True + ) + + def _strip_executor_escalation_instruction(self, text: str) -> str: + for snippet in (self._executor_system_instruction, EXECUTOR_INSTRUCTION): + if snippet and snippet in text: + text = text.replace(snippet, '') + return text.strip() + + async def _executor_instruction( + self, tool_context: ToolContext + ) -> str | None: + """Best-effort read of the executor agent's own instruction.""" + if not self._include_agent_instruction: + return None + invocation_context = tool_context._invocation_context + agent: Any = invocation_context.agent + canonical: Any = getattr(agent, 'canonical_instruction', None) + if not callable(canonical): + return None + + parts: list[str] = [] + static_inst: Any = getattr(agent, 'static_instruction', None) + if static_inst: + static_lines = [ + self._strip_executor_escalation_instruction(text) + for text in _extract_static_instruction_texts(static_inst) + ] + static_lines = [line for line in static_lines if line] + if static_lines: + parts.append('\n'.join(static_lines)) + + readonly_ctx = ReadonlyContext(invocation_context) + try: + instruction, bypass_state_injection = await canonical(readonly_ctx) + if instruction and not bypass_state_injection: + try: + instruction = await inject_session_state(instruction, readonly_ctx) + except Exception: # pylint: disable=broad-exception-caught + logger.debug( + 'ModelConsultTool could not inject session state into' + ' instruction', + exc_info=True, + ) + for key, val in invocation_context.session.state.items(): + if isinstance(key, str) and key.isidentifier(): + replacement = '' if val is None else str(val) + instruction = instruction.replace(f'{{{key}}}', replacement) + except Exception: # pylint: disable=broad-exception-caught + logger.debug( + 'ModelConsultTool could not read the agent instruction', exc_info=True + ) + return '\n\n'.join(parts) or None + instruction = self._strip_executor_escalation_instruction(instruction or '') + if instruction: + parts.append(instruction) + return '\n\n'.join(parts) or None + + async def _tool_inventory(self, tool_context: ToolContext) -> str | None: + """Lists the executor's other tools so guidance can name them.""" + if not self._include_tool_inventory: + return None + invocation_context = tool_context._invocation_context + tools = invocation_context.canonical_tools_cache + if tools is None: + agent: Any = invocation_context.agent + canonical_tools: Any = getattr(agent, 'canonical_tools', None) + if not callable(canonical_tools): + return None + try: + tools = await canonical_tools(ReadonlyContext(invocation_context)) + except Exception: # pylint: disable=broad-exception-caught + logger.debug( + "ModelConsultTool could not read the agent's tools", exc_info=True + ) + return None + invocation_context.canonical_tools_cache = tools + + lines: list[str] = [] + for tool in tools: + if tool.name == self.name: + continue + description = ' '.join((tool.description or '').split()) + if len(description) > _TOOL_DESCRIPTION_LIMIT: + description = description[:_TOOL_DESCRIPTION_LIMIT] + '...' + lines.append( + f'- {tool.name}: {description}' if description else f'- {tool.name}' + ) + return '\n'.join(lines) or None + + def _normalized_agent_name(self, tool_context: ToolContext) -> str | None: + raw_name: Any = tool_context.agent_name + if not isinstance(raw_name, str): + return None + name = raw_name.strip() + return name if name and name != 'unknown' else None + + def _system_instruction( + self, + executor_instruction: str | None, + tool_inventory: str | None, + agent_name: str | None, + ) -> str: + blocks = [self._advisor_instruction] + if executor_instruction: + agent_label = agent_name or 'the executor' + blocks.append( + '--- EXECUTOR AGENT INSTRUCTION' + f' ({agent_label}) ---\nThe executor operates under the following' + ' instruction. Your guidance must respect' + f' it.\n\n{executor_instruction}' + ) + if tool_inventory: + blocks.append( + '--- TOOLS AVAILABLE TO THE EXECUTOR ---\nThese are the only tools' + ' the executor can call. Name them explicitly in your plan, with' + ' concrete arguments. Do not propose steps that require tools not' + f' listed here.\n\n{tool_inventory}' + ) + return '\n\n'.join(blocks) + + def _handoff_content( + self, question: str, context: str | None, agent_name: str | None + ) -> types.Content: + context_block = ( + CONTEXT_BLOCK_TEMPLATE.format(context=context.strip()) + if context and context.strip() + else '' + ) + agent_clause = f' ({agent_name})' if agent_name else '' + text = ADVISOR_HANDOFF_TEMPLATE.format( + agent_clause=agent_clause, + question=question.strip(), + context_block=context_block, + ) + return types.Content(role='user', parts=[types.Part(text=text)]) + + def _in_flight_consult_call_ids(self, events: list[Event]) -> list[str]: + """Returns unanswered model_consult function call ids in the event log.""" + answered_ids: set[str] = set() + consult_call_ids: list[str] = [] + for event in events: + if event.content is None or not event.content.parts: + continue + for part in event.content.parts: + fr = part.function_response + if fr is not None and fr.id: + answered_ids.add(fr.id) + fc = part.function_call + if fc is not None and fc.name == self.name and fc.id: + consult_call_ids.append(fc.id) + return [cid for cid in consult_call_ids if cid not in answered_ids] + + def _build_contents( + self, tool_context: ToolContext, question: str, context: str | None + ) -> list[types.Content]: + events = list(tool_context.session.events) + in_flight_consult_ids = self._in_flight_consult_call_ids(events) + skip_ids = [ + call_id + for call_id in ( + tool_context.function_call_id, + *in_flight_consult_ids, + ) + if call_id + ] + session_contents = build_advisor_contents( + events, config=self._context_config, skip_function_call_ids=skip_ids + ) + agent_name = self._normalized_agent_name(tool_context) + handoff = self._handoff_content(question, context, agent_name) + + if self._context_config.mode == 'transcript' and session_contents: + transcript = render_transcript(session_contents) + session_contents = [ + types.Content( + role='user', + parts=[ + types.Part( + text=( + f'--- EXECUTOR SESSION TRANSCRIPT ---\n\n{transcript}' + ) + ) + ], + ) + ] + return _merge_adjacent([*session_contents, handoff]) + + async def run_async( + self, *, args: dict[str, Any], tool_context: ToolContext + ) -> dict[str, Any]: + """Consults the advisor and returns its guidance. + + Args: + args: Tool call arguments (`question` and optional `context`). + tool_context: The execution context for the current tool call. + + Returns: + A structured dictionary with `status` set to `'ok'`, `'limit_reached'`, + `'error'`, or `'invalid_request'`. Never raises: budget exhaustion and + advisor failures return a response the executor can read and continue + from. + """ + raw_question = args.get('question') + question = ( + raw_question.strip() + if isinstance(raw_question, str) + else str(raw_question or '').strip() + ) + if not question: + return { + 'status': 'invalid_request', + 'message': ( + '`question` is required: state the decision or blocker you want' + ' reviewed.' + ), + } + + session = tool_context.session + session_id_key = id(session) + inv_key = tool_context.invocation_id or 'unknown' + state = self._session_states.get(session_id_key) + if state is not None and state.session_ref is not None: + if state.session_ref() is not session: + state = None + if state is None: + state = _SessionConsultState(session_ref=weakref.ref(session)) + self._session_states[session_id_key] = state + weakref.finalize(session, self._session_states.pop, session_id_key, None) + state_delta = tool_context.actions.state_delta + reserved = False + + try: + async with state.cond: + state.active_calls += 1 + deltas_for_inv = state.inv_deltas.setdefault(inv_key, []) + if not any(existing is state_delta for existing in deltas_for_inv): + deltas_for_inv.append(state_delta) + + while True: + turn_uses = self._turn_uses_so_far(tool_context) + session_uses = self._session_uses_so_far(tool_context) + + if ( + self.session_max_uses is not None + and session_uses >= self.session_max_uses + ): + logger.info( + 'ModelConsultTool session budget exhausted (%s/%s)', + session_uses, + self.session_max_uses, + ) + return { + 'status': 'limit_reached', + 'message': _SESSION_LIMIT_MESSAGE.format( + session_max_uses=self.session_max_uses + ), + 'consults': self._consult_stats(turn_uses, session_uses), + } + + if self.max_uses is not None and turn_uses >= self.max_uses: + logger.info( + 'ModelConsultTool turn budget exhausted for invocation %s' + ' (%s/%s)', + tool_context.invocation_id or '?', + turn_uses, + self.max_uses, + ) + return { + 'status': 'limit_reached', + 'message': _TURN_LIMIT_MESSAGE.format(max_uses=self.max_uses), + 'consults': self._consult_stats(turn_uses, session_uses), + } + + turn_reserved = state.reserved_turn.get(inv_key, 0) + session_saturated = ( + self.session_max_uses is not None + and session_uses + state.reserved_session >= self.session_max_uses + ) + turn_saturated = ( + self.max_uses is not None + and turn_uses + turn_reserved >= self.max_uses + ) + if not session_saturated and not turn_saturated: + state.reserved_session += 1 + state.reserved_turn[inv_key] = turn_reserved + 1 + reserved = True + break + await state.cond.wait() + + raw_context = args.get('context') + extra_context = ( + raw_context.strip() + if isinstance(raw_context, str) + else (str(raw_context).strip() if raw_context is not None else None) + ) + contents = self._build_contents(tool_context, question, extra_context) + executor_instruction = await self._executor_instruction(tool_context) + tool_inventory = await self._tool_inventory(tool_context) + system_instruction = self._system_instruction( + executor_instruction, + tool_inventory, + self._normalized_agent_name(tool_context), + ) + + try: + result = await call_advisor( + self.advisor_model, + contents=contents, + system_instruction=system_instruction, + thinking_level=self._thinking_level, + generate_content_config=self._generate_content_config, + timeout_seconds=self._timeout_seconds, + agent_name=self.name, + ) + except AdvisorError as exc: + logger.warning('ModelConsultTool advisor call failed: %s', exc) + async with state.cond: + turn_uses = self._turn_uses_so_far(tool_context) + session_uses = self._session_uses_so_far(tool_context) + return { + 'status': 'error', + 'error': str(exc), + 'message': _ERROR_MESSAGE, + 'advisor_model': self.advisor_model.model, + 'consults': self._consult_stats(turn_uses, session_uses), + } + + async with state.cond: + turn_uses = self._turn_uses_so_far(tool_context) + 1 + session_uses = self._session_uses_so_far(tool_context) + 1 + self._record_use( + tool_context, + turn_uses=turn_uses, + session_uses=session_uses, + deltas=state.inv_deltas.get(inv_key, []), + ) + return self._success_payload(result, turn_uses, session_uses) + finally: + async with state.cond: + if reserved: + state.reserved_session = max(state.reserved_session - 1, 0) + remaining_reserved = state.reserved_turn.get(inv_key, 1) - 1 + if remaining_reserved <= 0: + state.reserved_turn.pop(inv_key, None) + else: + state.reserved_turn[inv_key] = remaining_reserved + state.cond.notify_all() + state.active_calls = max(state.active_calls - 1, 0) + if state.active_calls == 0: + state.inv_deltas.clear() + + def _consult_stats(self, turn_uses: int, session_uses: int) -> dict[str, Any]: + turn_remaining = ( + None if self.max_uses is None else max(self.max_uses - turn_uses, 0) + ) + session_remaining = ( + None + if self.session_max_uses is None + else max(self.session_max_uses - session_uses, 0) + ) + if turn_remaining is not None and session_remaining is not None: + remaining: int | None = min(turn_remaining, session_remaining) + elif turn_remaining is not None: + remaining = turn_remaining + else: + remaining = session_remaining + + return { + 'used_this_turn': turn_uses, + 'max_uses': self.max_uses, + 'used_this_session': session_uses, + 'session_max_uses': self.session_max_uses, + 'remaining': remaining, + } + + def _success_payload( + self, result: AdvisorResult, turn_uses: int, session_uses: int + ) -> dict[str, Any]: + return { + 'status': 'ok', + 'guidance': result.text, + 'advisor_model': result.model_version or result.model, + 'thinking_level': self.thinking_level, + 'consults': self._consult_stats(turn_uses, session_uses), + 'usage': result.usage.to_dict(), + 'latency_ms': result.latency_ms, + } diff --git a/src/google/adk/tools/model_consult/_prompts.py b/src/google/adk/tools/model_consult/_prompts.py new file mode 100644 index 00000000000..3944ea0fccf --- /dev/null +++ b/src/google/adk/tools/model_consult/_prompts.py @@ -0,0 +1,142 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Prompt assets for ModelConsultTool. + +Three prompt strings live here, all overridable on `ModelConsultTool`: + +* `TOOL_DESCRIPTION` -- the tool schema description read by the executor model + when deciding whether to escalate (override via `description`). +* `EXECUTOR_INSTRUCTION` -- escalation policy automatically appended to the + executor's `system_instruction` by `ModelConsultTool.process_llm_request` + (override via `executor_instruction`, or pass `""` to disable). +* `ADVISOR_SYSTEM_INSTRUCTION` -- the advisor's role. It must produce a plan + or a course correction, not the finished deliverable, so that the bulk of + token generation stays at executor rates (override via `advisor_instruction`). +""" + +from __future__ import annotations + +# Two empirical findings shape the timing guidance below: +# * A consult placed before the executor has gathered any context is +# low-value and can displace a better-timed later call -- hence the +# explicit carve-out that orientation is not substantive work. +# * A second consult before declaring done is worth about as much as the +# first, so the target cadence is two to three calls per task, not one. + +TOOL_DESCRIPTION = """\ +Consult a stronger advisor model for strategic guidance. The advisor sees this \ +entire conversation -- your instructions, your reasoning, every tool call you \ +made and every result you saw -- so do not restate the task. + +Call this tool BEFORE substantive work: before writing or editing, before \ +committing to an interpretation, before building on an assumption. If the task \ +needs orientation first (finding files, reading the issue, seeing what is \ +there), do that first, then call. Orientation is not substantive work. \ +Writing, editing and declaring an answer are. + +Also call this tool: +- When you believe the task is complete, before you declare it done. Make any \ +deliverable durable first (write the file, save the result). +- When stuck: errors recurring, an approach not converging, results that do \ +not fit. +- When considering a change of approach. + +On tasks longer than a few steps, call once before committing to an approach \ +and once before declaring done. On short reactive tasks where the next action \ +is dictated by output you just read, do not keep calling: most of the value is \ +in the first well-timed call. + +Returns a plan or course correction, not a finished answer. You still do the \ +work.\ +""" + +EXECUTOR_INSTRUCTION = """\ +You have access to a `model_consult` tool backed by a stronger advisor model. \ +It sees your entire conversation, so pass only the specific decision you want \ +reviewed. + +Call `model_consult` BEFORE substantive work: before writing, before \ +committing to an interpretation, before building on an assumption. If the task \ +needs orientation first (finding files, reading the issue, seeing what is \ +there), do that first, then call. Orientation is not substantive work. \ +Writing, editing and declaring an answer are. + +Also call `model_consult`: +- When you believe the task is complete, before declaring it done. Make your \ +deliverable durable first. +- When stuck: errors recurring, an approach not converging, results that do \ +not fit. +- When considering a change of approach. + +On tasks longer than a few steps, call at least once before committing to an \ +approach and once before declaring done. + +Give the advice serious weight. Adapt only if a step fails empirically or you \ +have primary-source evidence that contradicts a specific claim; a passing \ +self-check is not evidence the advice is wrong. If your own evidence points \ +one way and the advisor points another, do not silently switch: say what you \ +found, say what it suggested, and ask which constraint breaks the tie.\ +""" + +ADVISOR_SYSTEM_INSTRUCTION = """\ +You are a senior technical advisor consulted mid-task by a faster, smaller \ +executor agent. You are reading the executor's full working session: its \ +instructions, its reasoning so far, the tools it called and what those tools \ +returned. + +Your job is to make the executor's NEXT steps correct and efficient. Produce a \ +plan or a course correction -- not the finished deliverable. The executor does \ +the work; you decide what the work should be. + +Answer with: +1. Diagnosis -- in one or two sentences, what is actually going on, including \ +any mistaken assumption the executor is operating under. +2. Plan -- concrete numbered next steps the executor can act on directly. Name \ +specific tools, files, commands, identifiers and values wherever the session \ +gives you enough to be specific. Vague advice is worse than none. +3. Watch out for -- the failure modes, edge cases or verification steps most \ +likely to bite, and how the executor will know it is on the wrong track. + +Rules: +- Be concrete and brief. Aim for under 300 words; never pad. +- If the executor is already on the right track, say so plainly and give the \ +shortest path to done rather than inventing a new approach. +- If the session lacks information you need, say exactly what the executor \ +should gather and how, instead of guessing. +- Short code or command snippets are fine when they are the clearest way to \ +specify a step. Do not write out the whole solution. +- Never ask the executor a question back; it cannot reply. Decide, and state \ +the assumption you decided under.\ +""" + +# Framing appended as the final user turn of the advisor request. Keeps the +# advisor from simply continuing the conversation as if it were the executor. +ADVISOR_HANDOFF_TEMPLATE = """\ +--- END OF EXECUTOR SESSION --- + +You are now being consulted as the advisor. The executor agent{agent_clause} \ +paused its work and asked you: + +{question} +{context_block} +Respond with the diagnosis / plan / watch-out-for structure. Advise the \ +executor on its next steps; do not produce the final deliverable yourself.\ +""" + +CONTEXT_BLOCK_TEMPLATE = """ +Additional context the executor supplied: + +{context} +""" diff --git a/tests/unittests/tools/model_consult/test_model_consult_tool.py b/tests/unittests/tools/model_consult/test_model_consult_tool.py new file mode 100644 index 00000000000..7560a3000f7 --- /dev/null +++ b/tests/unittests/tools/model_consult/test_model_consult_tool.py @@ -0,0 +1,1354 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Unit tests for ModelConsultTool.""" + +from __future__ import annotations + +import asyncio +from collections.abc import AsyncGenerator +from typing import Any + +from google.adk.agents.invocation_context import InvocationContext +from google.adk.agents.llm_agent import LlmAgent +from google.adk.events.event import Event +from google.adk.events.event_actions import EventActions +from google.adk.flows.llm_flows.functions import merge_parallel_function_response_events +from google.adk.models.base_llm import BaseLlm +from google.adk.models.llm_request import LlmRequest +from google.adk.models.llm_response import LlmResponse +from google.adk.sessions.in_memory_session_service import InMemorySessionService +from google.adk.sessions.session import Session +from google.adk.sessions.state import State +from google.adk.tools import model_consult as model_consult_pkg +from google.adk.tools import ModelConsultContextConfig as TopLevelContextConfig +from google.adk.tools import ModelConsultTool as TopLevelModelConsultTool +from google.adk.tools.model_consult import ADVISOR_SYSTEM_INSTRUCTION +from google.adk.tools.model_consult import ContextMode +from google.adk.tools.model_consult import DEFAULT_ADVISOR_MODEL +from google.adk.tools.model_consult import DEFAULT_TOOL_NAME +from google.adk.tools.model_consult import EXECUTOR_INSTRUCTION +from google.adk.tools.model_consult import ModelConsultContextConfig +from google.adk.tools.model_consult import ModelConsultTool +from google.adk.tools.model_consult import TOOL_DESCRIPTION +from google.adk.tools.tool_context import ToolContext +from google.genai import types +from pydantic import BaseModel +from pydantic import Field +import pytest + + +def _text_response( + text: str = '1. Diagnosis. 2. Plan. 3. Watch out.', + *, + model_version: str | None = 'fake-advisor-001', + prompt_tokens: int = 1000, + output_tokens: int = 120, + thoughts_tokens: int = 50, +) -> LlmResponse: + return LlmResponse( + model_version=model_version, + content=types.Content( + role='model', + parts=[types.Part(text=text)], + ), + finish_reason=types.FinishReason.STOP, + usage_metadata=types.GenerateContentResponseUsageMetadata( + prompt_token_count=prompt_tokens, + candidates_token_count=output_tokens, + thoughts_token_count=thoughts_tokens, + total_token_count=prompt_tokens + output_tokens + thoughts_tokens, + ), + ) + + +class _FakeAdvisorLlm(BaseLlm): + """Deterministic in-memory advisor LLM for tool tests.""" + + model: str = 'fake-advisor' + responses: list[LlmResponse] = Field(default_factory=list) + errors: list[Exception | None] = Field(default_factory=list) + requests: list[LlmRequest] = Field(default_factory=list) + delay_seconds: float = 0.0 + per_call_delays: list[float] = Field(default_factory=list) + + async def generate_content_async( + self, llm_request: LlmRequest, stream: bool = False + ) -> AsyncGenerator[LlmResponse, None]: + del stream + self.requests.append(llm_request.model_copy(deep=True)) + call_idx = len(self.requests) - 1 + delay = ( + self.per_call_delays[call_idx] + if call_idx < len(self.per_call_delays) + else self.delay_seconds + ) + if delay > 0: + await asyncio.sleep(delay) + if call_idx < len(self.errors) and self.errors[call_idx] is not None: + raise self.errors[call_idx] + if not self.responses: + yield _text_response() + return + response = self.responses[min(call_idx, len(self.responses) - 1)] + yield response + + +def _user_event(text: str) -> Event: + return Event( + invocation_id='inv-1', + author='user', + content=types.Content(role='user', parts=[types.Part(text=text)]), + ) + + +def _agent_event(parts: list[types.Part], *, author: str = 'executor') -> Event: + return Event( + invocation_id='inv-1', + author=author, + content=types.Content(role='model', parts=parts), + ) + + +def _tool_result_event( + name: str, + response: dict[str, Any], + *, + call_id: str = 'fc-1', + author: str = 'executor', +) -> Event: + return Event( + invocation_id='inv-1', + author=author, + content=types.Content( + role='user', + parts=[ + types.Part( + function_response=types.FunctionResponse( + id=call_id, name=name, response=response + ) + ) + ], + ), + ) + + +def _make_tool_context( + events: list[Event] | None = None, + *, + instruction: str = 'Investigate production issues carefully.', + static_instruction: types.ContentUnion | None = None, + tools: list[Any] | None = None, + session: Session | None = None, + invocation_id: str = 'inv-1', + function_call_id: str | None = 'fc-consult', +) -> ToolContext: + agent = LlmAgent( + name='executor', + model='gemini-2.5-flash', + instruction=instruction, + static_instruction=static_instruction, + tools=tools or [], + ) + if session is None: + session = Session( + id='session-1', + app_name='test-app', + user_id='user-1', + state={}, + events=list(events or []), + ) + elif events is not None: + session.events = list(events) + invocation_context = InvocationContext( + session_service=InMemorySessionService(), + invocation_id=invocation_id, + agent=agent, + session=session, + ) + return ToolContext( + invocation_context, + function_call_id=function_call_id, + ) + + +async def _run( + tool: ModelConsultTool, tool_context: ToolContext, **args: Any +) -> dict[str, Any]: + return await tool.run_async(args=args, tool_context=tool_context) + + +def test_public_exports_and_prompt_constants(): + """Verifies public re-exports on tools and model_consult packages.""" + assert TopLevelModelConsultTool is ModelConsultTool + assert TopLevelContextConfig is ModelConsultContextConfig + expected_all = { + 'ADVISOR_SYSTEM_INSTRUCTION', + 'ContextMode', + 'DEFAULT_ADVISOR_MODEL', + 'DEFAULT_TOOL_NAME', + 'EXECUTOR_INSTRUCTION', + 'ModelConsultContextConfig', + 'ModelConsultTool', + 'TOOL_DESCRIPTION', + } + assert set(model_consult_pkg.__all__) == expected_all + assert ContextMode is not None + assert DEFAULT_TOOL_NAME == 'model_consult' + assert DEFAULT_ADVISOR_MODEL == 'gemini-3.1-pro-preview' + assert 'advisor' in TOOL_DESCRIPTION.lower() + assert '`model_consult`' in EXECUTOR_INSTRUCTION + assert 'senior technical advisor' in ADVISOR_SYSTEM_INSTRUCTION + + +def test_declaration_shape(): + """Verifies function declaration schema and required question field.""" + tool = ModelConsultTool(model=_FakeAdvisorLlm()) + + decl = tool._get_declaration() + + assert decl.name == 'model_consult' + assert decl.parameters is not None + assert decl.parameters.required == ['question'] + assert set(decl.parameters.properties or {}) == {'question', 'context'} + assert 'stuck' in (decl.description or '').lower() + + +def test_description_and_name_are_overridable(): + """Verifies custom name and description override defaults on declaration.""" + tool = ModelConsultTool( + model=_FakeAdvisorLlm(), + name='consult_expert', + description='Custom escalation description.', + ) + + assert tool.name == 'consult_expert' + assert tool._get_declaration().description == 'Custom escalation description.' + + +@pytest.mark.asyncio +async def test_process_llm_request_appends_executor_instruction_once(): + """Verifies process_llm_request injects EXECUTOR_INSTRUCTION without dupes.""" + tool = ModelConsultTool(model=_FakeAdvisorLlm()) + ctx = _make_tool_context([_user_event('go')]) + llm_request = LlmRequest() + llm_request.append_instructions(['You are an SRE assistant.']) + + await tool.process_llm_request(tool_context=ctx, llm_request=llm_request) + await tool.process_llm_request(tool_context=ctx, llm_request=llm_request) + + assert 'model_consult' in llm_request.tools_dict + sys_inst = llm_request.config.system_instruction or '' + assert sys_inst.count(EXECUTOR_INSTRUCTION) == 1 + + renamed_tool = ModelConsultTool(model=_FakeAdvisorLlm(), name='consult_sre') + renamed_request = LlmRequest() + await renamed_tool.process_llm_request( + tool_context=ctx, llm_request=renamed_request + ) + renamed_inst = renamed_request.config.system_instruction or '' + assert '`consult_sre`' in renamed_inst + assert '`model_consult`' not in renamed_inst + + custom_tool = ModelConsultTool( + model=_FakeAdvisorLlm(), + executor_instruction='Custom escalation rule.', + ) + custom_request = LlmRequest() + await custom_tool.process_llm_request( + tool_context=ctx, llm_request=custom_request + ) + assert ( + custom_request.config.system_instruction or '' + ) == 'Custom escalation rule.' + + disabled_tool = ModelConsultTool( + model=_FakeAdvisorLlm(), + executor_instruction='', + ) + disabled_request = LlmRequest() + await disabled_tool.process_llm_request( + tool_context=ctx, llm_request=disabled_request + ) + assert not (disabled_request.config.system_instruction or '') + + +@pytest.mark.asyncio +async def test_returns_guidance_and_accounting(): + """Verifies successful advisor consult returns guidance, usage, and budget.""" + llm = _FakeAdvisorLlm() + tool = ModelConsultTool(model=llm, max_uses=2, session_max_uses=5) + ctx = _make_tool_context([_user_event('Why is checkout slow?')]) + + result = await _run( + tool, ctx, question='Should I bisect deploys or profile CPU?' + ) + + assert result['status'] == 'ok' + assert result['guidance'] == '1. Diagnosis. 2. Plan. 3. Watch out.' + assert result['advisor_model'] == 'fake-advisor-001' + assert result['thinking_level'] == 'high' + assert result['consults'] == { + 'used_this_turn': 1, + 'max_uses': 2, + 'used_this_session': 1, + 'session_max_uses': 5, + 'remaining': 1, + } + assert result['usage']['prompt_tokens'] == 1000 + assert result['usage']['thoughts_tokens'] == 50 + assert result['latency_ms'] >= 0 + + +@pytest.mark.asyncio +async def test_advisor_sees_session_and_question(): + """Verifies session tool calls, tool results, and handoff reach advisor.""" + llm = _FakeAdvisorLlm() + tool = ModelConsultTool(model=llm) + ctx = _make_tool_context([ + _user_event('Investigate the paging alert.'), + _agent_event([ + types.Part( + function_call=types.FunctionCall( + id='fc-1', name='query_logs', args={'service': 'checkout'} + ) + ) + ]), + _tool_result_event('query_logs', {'errors': 42}), + ]) + + await _run( + tool, + ctx, + question='Which subsystem should I inspect next?', + context='p99 latency is flat across regions', + ) + + request = llm.requests[0] + texts = _extract_texts(request.contents) + assert 'Investigate the paging alert.' in texts + assert any('[tool_call] query_logs' in text for text in texts) + assert any( + '[tool_result] query_logs -> {"errors": 42}' in text for text in texts + ) + + handoff = request.contents[-1].parts[-1].text or '' + assert handoff.startswith('--- END OF EXECUTOR SESSION ---') + assert request.contents[-1].role == 'user' + assert ' (executor)' in handoff + assert 'Which subsystem should I inspect next?' in handoff + assert 'p99 latency is flat across regions' in handoff + + +def _extract_texts(contents: list[types.Content]) -> list[str]: + """Extracts all non-empty text strings from a list of Content messages.""" + texts: list[str] = [] + for content in contents: + for part in content.parts or []: + if part.text: + texts.append(part.text) + return texts + + +@pytest.mark.asyncio +async def test_contents_never_repeat_a_role(): + """Verifies adjacent turns in advisor contents strictly alternate roles.""" + llm = _FakeAdvisorLlm() + tool = ModelConsultTool(model=llm) + ctx = _make_tool_context([ + _user_event('first'), + _agent_event([types.Part(text='reply')]), + _tool_result_event('query_logs', {'errors': 1}), + ]) + + await _run(tool, ctx, question='Next?') + + roles = [content.role for content in llm.requests[0].contents] + assert all(left != right for left, right in zip(roles, roles[1:])) + + +@pytest.mark.asyncio +async def test_executor_instruction_is_forwarded_to_advisor(): + """Verifies executor instruction reaches advisor without escalation rules.""" + llm = _FakeAdvisorLlm() + tool = ModelConsultTool(model=llm, include_agent_instruction=True) + ctx = _make_tool_context( + [_user_event('go')], + instruction=( + f'Never restart production databases.\n\n{EXECUTOR_INSTRUCTION}' + ), + ) + + await _run(tool, ctx, question='Can I restart the DB?') + + system_inst = llm.requests[0].config.system_instruction + assert isinstance(system_inst, str) + assert 'senior technical advisor' in system_inst + assert 'Never restart production databases.' in system_inst + assert EXECUTOR_INSTRUCTION not in system_inst + + +@pytest.mark.asyncio +async def test_executor_instruction_withheld_when_disabled(): + """Verifies executor agent instruction is omitted when disabled.""" + llm = _FakeAdvisorLlm() + tool = ModelConsultTool(model=llm, include_agent_instruction=False) + ctx = _make_tool_context( + [_user_event('go')], instruction='Never restart production databases.' + ) + + await _run(tool, ctx, question='Can I restart the DB?') + + assert ( + 'Never restart production databases.' + not in llm.requests[0].config.system_instruction + ) + + +@pytest.mark.asyncio +async def test_executor_instruction_injects_state_and_static_instruction(): + """Verifies {state} placeholders and static_instruction reach the advisor.""" + llm = _FakeAdvisorLlm() + tool = ModelConsultTool(model=llm) + session = Session( + id='session-1', + app_name='app', + user_id='user-1', + state={'target_env': 'prod-eu-west'}, + events=[_user_event('go')], + ) + ctx = _make_tool_context( + session=session, + instruction='Only inspect cluster {target_env}.', + static_instruction=types.Content( + role='user', + parts=[types.Part(text='Global policy: read-only mode.')], + ), + ) + + await _run(tool, ctx, question='Which cluster?') + + system_inst = llm.requests[0].config.system_instruction + assert 'Global policy: read-only mode.' in system_inst + assert 'Only inspect cluster prod-eu-west.' in system_inst + + # Verify string static_instruction and fallback when an unset {placeholder} + # coexists with a populated {target_env} state key. + ctx_fallback = _make_tool_context( + session=session, + instruction='Cluster {target_env} with {unset_var}.', + static_instruction='String static instruction.', + invocation_id='inv-2', + ) + await _run(tool, ctx_fallback, question='Fallback check?') + system_inst_2 = llm.requests[1].config.system_instruction + assert 'String static instruction.' in system_inst_2 + assert 'Cluster prod-eu-west with {unset_var}.' in system_inst_2 + + # Verify callable instruction provider (bypass_state_injection=True) + ctx_provider = _make_tool_context( + session=session, + instruction=lambda _: 'Callable provider {target_env} literal.', + invocation_id='inv-3', + ) + await _run(tool, ctx_provider, question='Provider check?') + system_inst_3 = llm.requests[2].config.system_instruction + assert 'Callable provider {target_env} literal.' in system_inst_3 + + # Verify Part and list ContentUnion forms of static_instruction. + ctx_part = _make_tool_context( + session=session, + instruction='Dynamic instruction.', + static_instruction=types.Part(text='Part static instruction.'), + invocation_id='inv-4', + ) + await _run(tool, ctx_part, question='Part static check?') + system_inst_4 = llm.requests[3].config.system_instruction + assert 'Part static instruction.' in system_inst_4 + + ctx_list = _make_tool_context( + session=session, + instruction='Dynamic instruction.', + static_instruction=[ + 'List static part 1.', + types.Part(text='List static part 2.'), + {'text': 'Dict static part 3.'}, + types.Part.from_bytes(data=b'img', mime_type='image/png'), + types.File(uri='gs://bucket/doc.pdf'), + ], + invocation_id='inv-5', + ) + await _run(tool, ctx_list, question='List static check?') + system_inst_5 = llm.requests[4].config.system_instruction + assert ( + 'List static part 1.\nList static part 2.\nDict static part 3.' + in system_inst_5 + ) + + +@pytest.mark.asyncio +async def test_pending_model_consult_call_is_not_duplicated(): + """Verifies in-flight model_consult calls are skipped while completed stay.""" + llm = _FakeAdvisorLlm() + tool = ModelConsultTool(model=llm) + ctx = _make_tool_context( + [ + _user_event('go'), + Event( + author='executor', + content=types.Content(role='model', parts=[]), + ), + _agent_event([ + types.Part( + function_call=types.FunctionCall( + id='fc-answered', + name='model_consult', + args={'question': 'Earlier question?'}, + ) + ) + ]), + Event( + author='executor', + content=types.Content( + role='user', + parts=[ + types.Part( + function_response=types.FunctionResponse( + id='fc-answered', + name='model_consult', + response={'guidance': 'Check connection pool.'}, + ) + ) + ], + ), + ), + _agent_event([ + types.Part( + function_call=types.FunctionCall( + id='fc-current', + name='model_consult', + args={'question': 'What now?'}, + ) + ), + types.Part( + function_call=types.FunctionCall( + id='fc-sibling-parallel', + name='model_consult', + args={'question': 'Parallel question?'}, + ) + ), + ]), + ], + function_call_id='fc-current', + ) + + await _run(tool, ctx, question='What now?') + + texts = _extract_texts(llm.requests[0].contents) + assert any('Earlier question?' in text for text in texts) + assert any('Check connection pool.' in text for text in texts) + assert not any('Parallel question?' in text for text in texts) + assert not any( + '[tool_call] model_consult' in text and 'What now?' in text + for text in texts + ) + + +@pytest.mark.asyncio +async def test_transcript_mode_folds_session_into_one_turn(): + """Verifies transcript mode collapses session into a single user Content.""" + llm = _FakeAdvisorLlm() + tool = ModelConsultTool( + model=llm, + context_config=ModelConsultContextConfig(mode='transcript'), + ) + ctx = _make_tool_context([ + _user_event('go'), + _agent_event([types.Part(text='checking logs')]), + ]) + + await _run(tool, ctx, question='Next?') + + contents = llm.requests[0].contents + assert len(contents) == 1 + assert contents[0].role == 'user' + assert len(contents[0].parts) == 2 + assert 'EXECUTOR SESSION TRANSCRIPT' in (contents[0].parts[0].text or '') + assert 'USER: go' in (contents[0].parts[0].text or '') + assert 'AGENT: checking logs' in (contents[0].parts[0].text or '') + assert (contents[0].parts[-1].text or '').startswith( + '--- END OF EXECUTOR SESSION ---' + ) + + +@pytest.mark.asyncio +async def test_transcript_mode_does_not_charge_media_bytes_against_max_chars(): + """Verifies transcript mode converts media to text before char budgeting.""" + llm = _FakeAdvisorLlm() + tool = ModelConsultTool( + model=llm, + context_config=ModelConsultContextConfig( + mode='transcript', max_chars=500, include_media=True + ), + ) + ctx = _make_tool_context([ + _user_event('Initial root cause clue'), + _agent_event([types.Part(text='Middle investigation note')]), + Event( + invocation_id='inv-1', + author='user', + content=types.Content( + role='user', + parts=[ + types.Part(text='Screenshot attached'), + types.Part( + inline_data=types.Blob( + mime_type='image/png', data=b'x' * 10_000 + ) + ), + ], + ), + ), + ]) + + await _run(tool, ctx, question='Next?') + + transcript_part = llm.requests[0].contents[0].parts[0].text or '' + assert 'Middle investigation note' in transcript_part + assert '[media: image/png' in transcript_part + + +@pytest.mark.parametrize( + 'level,expected_enum,expected_name', + [ + ('minimal', types.ThinkingLevel.MINIMAL, 'minimal'), + ('low', types.ThinkingLevel.LOW, 'low'), + ('medium', types.ThinkingLevel.MEDIUM, 'medium'), + ('high', types.ThinkingLevel.HIGH, 'high'), + (types.ThinkingLevel.HIGH, types.ThinkingLevel.HIGH, 'high'), + ], +) +@pytest.mark.asyncio +async def test_thinking_level_reaches_request( + level: str | types.ThinkingLevel, + expected_enum: types.ThinkingLevel, + expected_name: str, +): + """Verifies string and enum thinking levels populate ThinkingConfig.""" + llm = _FakeAdvisorLlm() + tool = ModelConsultTool(model=llm, thinking_level=level) + ctx = _make_tool_context([_user_event('go')]) + + result = await _run(tool, ctx, question='Next?') + + assert llm.requests[0].config.thinking_config.thinking_level == expected_enum + assert result['thinking_level'] == expected_name + + +@pytest.mark.asyncio +async def test_thinking_level_none_sends_no_thinking_config(): + """Verifies thinking_level=None omits ThinkingConfig from advisor request.""" + llm = _FakeAdvisorLlm() + tool = ModelConsultTool(model=llm, thinking_level=None) + ctx = _make_tool_context([_user_event('go')]) + + result = await _run(tool, ctx, question='Next?') + + assert llm.requests[0].config.thinking_config is None + assert result['thinking_level'] is None + + +@pytest.mark.parametrize( + 'kwargs,error_match', + [ + ({'thinking_level': 'turbo'}, 'thinking_level'), + ({'max_uses': 0}, 'max_uses'), + ({'max_uses': -1}, 'max_uses'), + ({'session_max_uses': 0}, 'session_max_uses'), + ({'session_max_uses': -2}, 'session_max_uses'), + ({'max_output_tokens': 0}, 'max_output_tokens'), + ( + { + 'generate_content_config': types.GenerateContentConfig( + max_output_tokens=0 + ) + }, + 'generate_content_config.max_output_tokens', + ), + ( + { + 'max_output_tokens': 2048, + 'generate_content_config': types.GenerateContentConfig( + max_output_tokens=512 + ), + }, + 'Conflicting max_output_tokens', + ), + ({'timeout_seconds': 0}, 'timeout_seconds'), + ({'model': ' '}, 'non-empty model string'), + ], +) +def test_invalid_init_arguments_rejected_at_construction( + kwargs: dict[str, Any], error_match: str +): + """Verifies invalid init parameters raise ValueError at construction.""" + init_kwargs: dict[str, Any] = {'model': _FakeAdvisorLlm(), **kwargs} + with pytest.raises(ValueError, match=error_match): + ModelConsultTool(**init_kwargs) + + +def test_model_string_resolves_through_adk_registry(): + """Verifies model string resolves to a BaseLlm via LLMRegistry.""" + tool = ModelConsultTool(model='gemini-3.1-pro-preview') + + assert tool.advisor_model.model == 'gemini-3.1-pro-preview' + assert type(tool.advisor_model).__name__ == 'Gemini' + + +@pytest.mark.asyncio +async def test_multiple_tool_instances_have_independent_budgets(): + """Verifies distinct ModelConsultTool names track separate use budgets.""" + llm = _FakeAdvisorLlm() + arch_tool = ModelConsultTool( + model=llm, name='consult_arch', max_uses=1, session_max_uses=1 + ) + sec_tool = ModelConsultTool( + model=llm, name='consult_sec', max_uses=1, session_max_uses=1 + ) + session = Session( + id='session-1', app_name='app', user_id='user-1', state={}, events=[] + ) + ctx = _make_tool_context( + [_user_event('review design')], session=session, invocation_id='inv-1' + ) + + r_arch_1 = await _run(arch_tool, ctx, question='Check architecture') + r_arch_2 = await _run(arch_tool, ctx, question='Check architecture again') + r_sec_1 = await _run(sec_tool, ctx, question='Check security') + + assert r_arch_1['status'] == 'ok' + assert r_arch_2['status'] == 'limit_reached' + assert r_sec_1['status'] == 'ok' + + +@pytest.mark.asyncio +async def test_generate_content_config_does_not_mutate_input(): + """Verifies caller config is not mutated and max_output_tokens syncs.""" + llm = _FakeAdvisorLlm() + caller_cfg = types.GenerateContentConfig(temperature=0.2) + tool = ModelConsultTool( + model=llm, + max_output_tokens=2048, + generate_content_config=caller_cfg, + ) + ctx = _make_tool_context([_user_event('go')]) + + await _run(tool, ctx, question='Next?') + + sent_cfg = llm.requests[0].config + assert sent_cfg.temperature == 0.2 + assert sent_cfg.max_output_tokens == 2048 + assert tool.max_output_tokens == 2048 + assert caller_cfg.max_output_tokens is None + assert sent_cfg.system_instruction + + cfg_with_tokens = types.GenerateContentConfig( + temperature=0.3, max_output_tokens=512 + ) + tool_from_cfg = ModelConsultTool( + model=llm, + generate_content_config=cfg_with_tokens, + ) + assert tool_from_cfg.max_output_tokens == 512 + await _run(tool_from_cfg, ctx, question='Second?') + assert llm.requests[1].config.max_output_tokens == 512 + + +@pytest.mark.asyncio +async def test_max_uses_enforced_per_turn_and_resets_next_turn(): + """Verifies turn max_uses blocks excess calls and resets on next turn.""" + llm = _FakeAdvisorLlm() + tool = ModelConsultTool(model=llm, max_uses=1) + session = Session( + id='session-1', app_name='app', user_id='user-1', state={}, events=[] + ) + + turn1 = _make_tool_context( + [_user_event('turn 1')], session=session, invocation_id='inv-1' + ) + assert tool.has_remaining_budget(turn1) is True + first = await _run(tool, turn1, question='q1') + assert tool.has_remaining_budget(turn1) is False + second = await _run(tool, turn1, question='q2') + + assert first['status'] == 'ok' + assert second['status'] == 'limit_reached' + assert 'for this turn is exhausted (1 of 1 used)' in second['message'] + assert second['consults'] == { + 'used_this_turn': 1, + 'max_uses': 1, + 'used_this_session': 1, + 'session_max_uses': None, + 'remaining': 0, + } + assert len(llm.requests) == 1 + + turn2 = _make_tool_context( + [_user_event('turn 2')], session=session, invocation_id='inv-2' + ) + assert tool.has_remaining_budget(turn2) is True + third = await _run(tool, turn2, question='q3') + assert third['status'] == 'ok' + assert third['consults']['used_this_turn'] == 1 + assert third['consults']['used_this_session'] == 2 + assert len(llm.requests) == 2 + + +@pytest.mark.asyncio +async def test_session_max_uses_enforced_across_turns(): + """Verifies session_max_uses persists across turns and blocks once reached.""" + llm = _FakeAdvisorLlm() + tool = ModelConsultTool(model=llm, max_uses=2, session_max_uses=2) + session = Session( + id='session-1', app_name='app', user_id='user-1', state={}, events=[] + ) + + turn1 = _make_tool_context( + [_user_event('turn 1')], session=session, invocation_id='inv-1' + ) + r1 = await _run(tool, turn1, question='q1') + assert r1['status'] == 'ok' + assert r1['consults']['remaining'] == 1 + + turn2 = _make_tool_context( + [_user_event('turn 2')], session=session, invocation_id='inv-2' + ) + r2 = await _run(tool, turn2, question='q2') + assert r2['status'] == 'ok' + assert r2['consults']['remaining'] == 0 + + # Third turn has a fresh turn budget (0/2), but session budget (2/2) is full. + turn3 = _make_tool_context( + [_user_event('turn 3')], session=session, invocation_id='inv-3' + ) + assert tool.has_remaining_budget(turn3) is False + r3 = await _run(tool, turn3, question='q3') + assert r3['status'] == 'limit_reached' + assert 'for this session is exhausted (2 of 2 used)' in r3['message'] + assert r3['consults'] == { + 'used_this_turn': 0, + 'max_uses': 2, + 'used_this_session': 2, + 'session_max_uses': 2, + 'remaining': 0, + } + assert len(llm.requests) == 2 + assert session.state['model_consult:model_consult:session_uses'] == 2 + + +@pytest.mark.asyncio +async def test_session_max_uses_without_turn_cap_and_standalone_token_cap(): + """Verifies session_max_uses when max_uses is None and standalone cap.""" + llm = _FakeAdvisorLlm(responses=[_text_response(model_version=None)]) + tool = ModelConsultTool( + model=llm, + max_uses=None, + session_max_uses=2, + max_output_tokens=1024, + ) + ctx = _make_tool_context([_user_event('turn 1')]) + + r1 = await _run(tool, ctx, question='q1') + + assert r1['status'] == 'ok' + assert r1['advisor_model'] == 'fake-advisor' + assert llm.requests[0].config.max_output_tokens == 1024 + assert r1['consults'] == { + 'used_this_turn': 1, + 'max_uses': None, + 'used_this_session': 1, + 'session_max_uses': 2, + 'remaining': 1, + } + + +@pytest.mark.asyncio +async def test_parallel_consult_calls_respect_caps_and_preserve_deltas(): + """Verifies parallel model_consult calls serialize budgets and state_delta.""" + # 1) Turn cap saturation only (max_uses=1, session_max_uses=5). + llm_turn_cap = _FakeAdvisorLlm(delay_seconds=0.02) + tool_turn_cap = ModelConsultTool( + model=llm_turn_cap, max_uses=1, session_max_uses=5 + ) + session_turn = Session( + id='s-turn', app_name='app', user_id='u1', state={}, events=[] + ) + ctx_turn_a = _make_tool_context( + [_user_event('go')], + session=session_turn, + invocation_id='inv-turn', + function_call_id='fc-a', + ) + ctx_turn_b = _make_tool_context( + [_user_event('go')], + session=session_turn, + invocation_id='inv-turn', + function_call_id='fc-b', + ) + res_ta, res_tb = await asyncio.gather( + _run(tool_turn_cap, ctx_turn_a, question='q1'), + _run(tool_turn_cap, ctx_turn_b, question='q2'), + ) + assert sorted([res_ta['status'], res_tb['status']]) == ['limit_reached', 'ok'] + assert len(llm_turn_cap.requests) == 1 + + # 2) Session cap saturation only (max_uses=5, session_max_uses=1). + llm_sess_cap = _FakeAdvisorLlm(delay_seconds=0.02) + tool_sess_cap = ModelConsultTool( + model=llm_sess_cap, max_uses=5, session_max_uses=1 + ) + session_sess = Session( + id='s-sess', app_name='app', user_id='u1', state={}, events=[] + ) + ctx_sess_a = _make_tool_context( + [_user_event('go')], + session=session_sess, + invocation_id='inv-sess', + function_call_id='fc-sa', + ) + ctx_sess_b = _make_tool_context( + [_user_event('go')], + session=session_sess, + invocation_id='inv-sess', + function_call_id='fc-sb', + ) + res_sa, res_sb = await asyncio.gather( + _run(tool_sess_cap, ctx_sess_a, question='q1'), + _run(tool_sess_cap, ctx_sess_b, question='q2'), + ) + assert sorted([res_sa['status'], res_sb['status']]) == ['limit_reached', 'ok'] + assert len(llm_sess_cap.requests) == 1 + + # Now test max_uses=5 where Call 1 takes longer than Call 2 so Call 2 finishes + # first, and verify merge_parallel_function_response_events preserves count=2. + llm_cap5 = _FakeAdvisorLlm(per_call_delays=[0.03, 0.005, 0.03, 0.005]) + tool_cap5 = ModelConsultTool(model=llm_cap5, max_uses=5, session_max_uses=5) + session_service = InMemorySessionService() + session_cap5 = await session_service.create_session( + app_name='app', user_id='u1', session_id='s5' + ) + inv_ctx = InvocationContext( + session_service=session_service, + invocation_id='inv-5', + agent=LlmAgent(name='executor', model='gemini-2.5-flash'), + session=session_cap5, + ) + ctx5_1 = ToolContext( + inv_ctx, function_call_id='fc-1', event_actions=EventActions() + ) + ctx5_2 = ToolContext( + inv_ctx, function_call_id='fc-2', event_actions=EventActions() + ) + + r5_1, r5_2 = await asyncio.gather( + _run(tool_cap5, ctx5_1, question='q1'), + _run(tool_cap5, ctx5_2, question='q2'), + ) + ev1 = Event( + invocation_id='inv-5', + author='executor', + content=types.Content( + role='user', + parts=[ + types.Part.from_function_response( + name='model_consult', response=r5_1 + ) + ], + ), + actions=ctx5_1.actions, + ) + ev2 = Event( + invocation_id='inv-5', + author='executor', + content=types.Content( + role='user', + parts=[ + types.Part.from_function_response( + name='model_consult', response=r5_2 + ) + ], + ), + actions=ctx5_2.actions, + ) + merged_event = merge_parallel_function_response_events([ev1, ev2]) + await session_service.append_event(session=session_cap5, event=merged_event) + + assert session_cap5.state['model_consult:model_consult:session_uses'] == 2 + assert session_cap5.state['temp:model_consult:model_consult:inv-5:uses'] == 2 + + # Also verify reverse completion order (when fc-2 finishes before fc-1) still + # merges state_delta to 2 rather than overwriting 2 back to 1. + session_rev = await session_service.create_session( + app_name='app', user_id='u1' + ) + ctx_rev_1 = _make_tool_context( + [_user_event('rev')], + session=session_rev, + invocation_id='inv-rev', + function_call_id='fc-rev-1', + ) + ctx_rev_2 = _make_tool_context( + [_user_event('rev')], + session=session_rev, + invocation_id='inv-rev', + function_call_id='fc-rev-2', + ) + + r_rev_1, r_rev_2 = await asyncio.gather( + _run(tool_cap5, ctx_rev_1, question='rev-1'), + _run(tool_cap5, ctx_rev_2, question='rev-2'), + ) + ev_rev_1 = Event( + invocation_id='inv-rev', + author='executor', + content=types.Content( + role='user', + parts=[ + types.Part.from_function_response( + name='model_consult', response=r_rev_1 + ) + ], + ), + actions=ctx_rev_1.actions, + ) + ev_rev_2 = Event( + invocation_id='inv-rev', + author='executor', + content=types.Content( + role='user', + parts=[ + types.Part.from_function_response( + name='model_consult', response=r_rev_2 + ) + ], + ), + actions=ctx_rev_2.actions, + ) + merged_rev = merge_parallel_function_response_events([ev_rev_1, ev_rev_2]) + await session_service.append_event(session=session_rev, event=merged_rev) + assert session_rev.state['model_consult:model_consult:session_uses'] == 2 + assert session_rev.state['temp:model_consult:model_consult:inv-rev:uses'] == 2 + + # Verify two sequential consults in the same invocation do not mutate the + # already-emitted first event's state_delta (inv_deltas is pruned when + # active_calls drops to 0). + session_seq = await session_service.create_session( + app_name='app', user_id='u1' + ) + ctx_seq_1 = _make_tool_context( + [_user_event('seq')], + session=session_seq, + invocation_id='inv-seq', + function_call_id='fc-seq-1', + ) + ctx_seq_2 = _make_tool_context( + [_user_event('seq')], + session=session_seq, + invocation_id='inv-seq', + function_call_id='fc-seq-2', + ) + await _run(tool_cap5, ctx_seq_1, question='seq-1') + assert ( + ctx_seq_1.actions.state_delta[ + 'temp:model_consult:model_consult:inv-seq:uses' + ] + == 1 + ) + await _run(tool_cap5, ctx_seq_2, question='seq-2') + assert ( + ctx_seq_1.actions.state_delta[ + 'temp:model_consult:model_consult:inv-seq:uses' + ] + == 1 + ) + assert ( + ctx_seq_2.actions.state_delta[ + 'temp:model_consult:model_consult:inv-seq:uses' + ] + == 2 + ) + + +@pytest.mark.asyncio +async def test_session_max_uses_persists_with_strict_state_schema(): + """Verifies session_max_uses works even when State enforces a state_schema.""" + + class _StrictSchema(BaseModel): + allowed_field: str = 'ok' + + llm = _FakeAdvisorLlm() + tool = ModelConsultTool(model=llm, session_max_uses=1) + session = Session( + id='session-strict', app_name='app', user_id='u1', state={}, events=[] + ) + ctx1 = _make_tool_context( + [_user_event('t1')], session=session, invocation_id='inv-1' + ) + ctx1._state = State( + value=session.state, + delta=ctx1.actions.state_delta, + schema=_StrictSchema, + ) + + r1 = await _run(tool, ctx1, question='q1') + assert r1['status'] == 'ok' + assert session.state['model_consult:model_consult:session_uses'] == 1 + assert ( + ctx1.actions.state_delta['model_consult:model_consult:session_uses'] == 1 + ) + + ctx2 = _make_tool_context( + [_user_event('t2')], session=session, invocation_id='inv-2' + ) + ctx2._state = State( + value=session.state, + delta=ctx2.actions.state_delta, + schema=_StrictSchema, + ) + r2 = await _run(tool, ctx2, question='q2') + assert r2['status'] == 'limit_reached' + assert len(llm.requests) == 1 + + +@pytest.mark.asyncio +async def test_missing_question_rejected_without_calling_advisor(): + """Verifies blank or missing question returns invalid_request immediately.""" + llm = _FakeAdvisorLlm() + tool = ModelConsultTool(model=llm) + ctx = _make_tool_context([_user_event('go')]) + + result_blank = await _run(tool, ctx, question=' ') + result_missing = await tool.run_async(args={}, tool_context=ctx) + + assert result_blank['status'] == 'invalid_request' + assert result_missing['status'] == 'invalid_request' + assert llm.requests == [] + + +@pytest.mark.asyncio +async def test_advisor_failure_degrades_gracefully_without_burning_budget(): + """Verifies advisor runtime error returns status='error' and keeps budget.""" + failing_llm = _FakeAdvisorLlm( + errors=[RuntimeError('503 backend unavailable')] + ) + tool = ModelConsultTool(model=failing_llm, max_uses=1, session_max_uses=1) + ctx = _make_tool_context([_user_event('go')]) + + result = await _run(tool, ctx, question='Next?') + + assert result['status'] == 'error' + assert '503' in result['error'] + assert 'own best judgment' in result['message'] + assert result['consults']['used_this_turn'] == 0 + assert result['consults']['used_this_session'] == 0 + assert result['consults']['remaining'] == 1 + + +@pytest.mark.asyncio +async def test_advisor_timeout_degrades_gracefully_without_burning_budget(): + """Verifies advisor timeout returns status='error' and keeps budget.""" + slow_llm = _FakeAdvisorLlm(delay_seconds=0.2) + timeout_tool = ModelConsultTool( + model=slow_llm, max_uses=1, session_max_uses=1, timeout_seconds=0.01 + ) + ctx = _make_tool_context([_user_event('go')]) + + timeout_result = await _run(timeout_tool, ctx, question='Next?') + + assert timeout_result['status'] == 'error' + assert 'timed out' in timeout_result['error'] + assert timeout_result['consults']['used_this_turn'] == 0 + assert timeout_result['consults']['remaining'] == 1 + + +@pytest.mark.asyncio +async def test_thinking_config_rejection_falls_back_and_still_answers(): + """Verifies unsupported thinking_level falls back without thinking_config.""" + llm = _FakeAdvisorLlm( + responses=[_text_response('fallback advice')], + errors=[ + ValueError('thinking_level is not supported by this model'), + None, + ], + ) + tool = ModelConsultTool(model=llm, thinking_level='high') + ctx = _make_tool_context([_user_event('go')]) + + result = await _run(tool, ctx, question='Next?') + + assert result['status'] == 'ok' + assert result['guidance'] == 'fallback advice' + assert len(llm.requests) == 2 + assert llm.requests[0].config.thinking_config is not None + assert llm.requests[1].config.thinking_config is None + + +@pytest.mark.asyncio +async def test_advisor_receives_executor_tool_inventory(): + """Verifies executor tools and truncated descriptions reach advisor prompt.""" + + def list_deploys(service: str) -> dict[str, str]: + """Lists recent deploys for a service.""" + return {'service': service} + + def verbose_tool(query: str) -> str: + return query + + verbose_tool.__doc__ = 'A' * 350 + + def no_doc_tool(x: str) -> str: + return x + + llm = _FakeAdvisorLlm() + tool = ModelConsultTool( + model=llm, + advisor_instruction='Custom advisor system prompt.', + max_uses=2, + ) + ctx = _make_tool_context( + [_user_event('go')], + tools=[list_deploys, verbose_tool, no_doc_tool, tool], + ) + assert ctx._invocation_context.canonical_tools_cache is None + + await _run(tool, ctx, question='What next?') + + assert ctx._invocation_context.canonical_tools_cache is not None + system = llm.requests[0].config.system_instruction + assert isinstance(system, str) + assert system.startswith('Custom advisor system prompt.') + assert 'TOOLS AVAILABLE TO THE EXECUTOR' in system + inventory_section = system.split('TOOLS AVAILABLE TO THE EXECUTOR')[1] + assert ( + '- list_deploys: Lists recent deploys for a service.' in inventory_section + ) + assert f"- verbose_tool: {'A' * 300}..." in inventory_section + assert '- no_doc_tool' in inventory_section + assert '- no_doc_tool:' not in inventory_section + assert 'model_consult' not in inventory_section + + # Second call within the same invocation reuses canonical_tools_cache + # without calling agent.canonical_tools again. + async def _fail_if_called(_): + raise AssertionError('canonical_tools should not be re-resolved') + + object.__setattr__( + ctx._invocation_context.agent, 'canonical_tools', _fail_if_called + ) + await _run(tool, ctx, question='Second check?') + system_2 = llm.requests[1].config.system_instruction + assert '- list_deploys: Lists recent deploys for a service.' in system_2 + + +@pytest.mark.asyncio +async def test_tool_inventory_withheld_when_disabled(): + """Verifies include_tool_inventory=False omits tool list from prompt.""" + + def list_deploys(service: str) -> dict[str, str]: + """Lists recent deploys for a service.""" + return {'service': service} + + llm = _FakeAdvisorLlm() + tool = ModelConsultTool(model=llm, include_tool_inventory=False) + ctx = _make_tool_context([_user_event('go')], tools=[list_deploys, tool]) + + await _run(tool, ctx, question='What next?') + + assert ( + 'TOOLS AVAILABLE TO THE EXECUTOR' + not in llm.requests[0].config.system_instruction + ) + + +@pytest.mark.asyncio +async def test_corrupt_state_and_broken_agent_callbacks_degrade_gracefully( + caplog: pytest.LogCaptureFixture, +): + """Verifies corrupt state counters and broken callbacks do not crash.""" + llm = _FakeAdvisorLlm() + tool = ModelConsultTool(model=llm, max_uses=3) + ctx = _make_tool_context( + [_user_event('go')], + instruction='Executor rule.', + invocation_id='', + ) + ctx.state[tool._turn_uses_state_key(ctx)] = -5 + ctx.state[tool._session_uses_state_key()] = 'not-an-int' + object.__setattr__(ctx._invocation_context.agent, 'name', 123) + + res0 = await _run(tool, ctx, question='Non-str agent name check?') + assert res0['status'] == 'ok' + assert '(the executor)' in llm.requests[0].config.system_instruction + assert tool._turn_uses_state_key(ctx).endswith(':unknown:uses') + + ctx._invocation_context.agent.name = 'unknown' + + async def _broken_instruction(_): + raise RuntimeError('instruction callback boom') + + async def _broken_tools(_): + raise RuntimeError('tools callback boom') + + ctx._invocation_context.canonical_tools_cache = None + object.__setattr__( + ctx._invocation_context.agent, + 'canonical_instruction', + _broken_instruction, + ) + object.__setattr__( + ctx._invocation_context.agent, + 'canonical_tools', + _broken_tools, + ) + + res = await _run(tool, ctx, question=12345, context=67890) + + assert res['status'] == 'ok' + assert res['consults']['used_this_turn'] == 2 + assert res['consults']['used_this_session'] == 2 + handoff_text = llm.requests[1].contents[-1].parts[-1].text or '' + assert '(unknown)' not in handoff_text + assert '12345' in handoff_text + assert '67890' in handoff_text + + class _RaisingStateDict(dict): + """State mapping that raises RuntimeError on write.""" + + def __setitem__(self, key, value): + raise RuntimeError('storage write failure') + + # Verify non-callable instruction/tools attributes and failing state write. + ctx._invocation_context.canonical_tools_cache = None + object.__setattr__(ctx._invocation_context.agent, 'canonical_instruction', 42) + object.__setattr__(ctx._invocation_context.agent, 'canonical_tools', 42) + object.__setattr__(ctx, '_state', _RaisingStateDict()) + caplog.clear() + res2 = await _run(tool, ctx, question='Still works?') + assert res2['status'] == 'ok' + assert any( + record.levelname == 'WARNING' + and 'ModelConsultTool could not persist its use counters' + in record.getMessage() + for record in caplog.records + ) From 89ebdecf233b9d57fcd27cc4b003a9a62b107097 Mon Sep 17 00:00:00 2001 From: Xuan Yang Date: Tue, 29 Sep 2026 17:01:49 -0700 Subject: [PATCH 26/29] docs(tools): add ModelConsultTool developer guide and sample agent Adds the developer guide and runnable order-support refund policy sample for `ModelConsultTool`: - Adds the `ModelConsultTool` unit guide under `docs/guides/tools/model_consult/model_consult_tool/index.md` covering getting started, how mid-generation escalation works, `ModelConsultTool` and `ModelConsultContextConfig` options, custom context budgets, custom `BaseLlm` advisors, and limitations, and registers it in `docs/guides/README.md`. - Adds a runnable e-commerce order support sample under `contributing/samples/tools/model_consult/` demonstrating `get_order`, `get_customer_profile`, and `issue_refund` paired with `ModelConsultTool`. Co-authored-by: Xuan Yang PiperOrigin-RevId: 990615455 --- .../samples/tools/model_consult/README.md | 131 ++++++++++++ .../samples/tools/model_consult/__init__.py | 15 ++ .../samples/tools/model_consult/agent.py | 157 ++++++++++++++ docs/guides/README.md | 1 + .../model_consult/model_consult_tool/index.md | 195 ++++++++++++++++++ 5 files changed, 499 insertions(+) create mode 100644 contributing/samples/tools/model_consult/README.md create mode 100644 contributing/samples/tools/model_consult/__init__.py create mode 100644 contributing/samples/tools/model_consult/agent.py create mode 100644 docs/guides/tools/model_consult/model_consult_tool/index.md diff --git a/contributing/samples/tools/model_consult/README.md b/contributing/samples/tools/model_consult/README.md new file mode 100644 index 00000000000..740840b88ea --- /dev/null +++ b/contributing/samples/tools/model_consult/README.md @@ -0,0 +1,131 @@ +# ADK Model Consult Sample + +## Overview + +This sample demonstrates how an e-commerce order support assistant, `order_support_agent`, pairs routine lookup and action tools, `get_order`, `get_customer_profile`, and `issue_refund`, with `ModelConsultTool` to escalate multi-rule refund policy decisions to a stronger advisor model mid-generation. + +The primary agent gathers order and customer details directly and follows the default escalation policy that `ModelConsultTool` adds to its system instruction: it calls `model_consult` before committing to a refund decision and, on longer tasks, again before declaring the task done. The advisor adds the most value when multiple policy exceptions interact, such as late returns, opened electronics restocking fees, defect bulletins, and Gold-tier loyalty exemptions. The agent then executes `issue_refund` based on the advisor's guidance. + +## Sample Inputs + +- `Customer CUST-108 wants a full refund to their original payment method for order ORD-502 (wireless headphones bought 45 days ago, opened, battery drains quickly). Check the order and customer profile, process the appropriate refund, and explain the decision.` + + *The agent calls `get_order('ORD-502')` and `get_customer_profile('CUST-108')`, consults `model_consult` to reconcile the 30-day return cutoff against defect bulletin `SB-2026-04` and the customer's Gold-tier loyalty status with a `2.4%` return rate, executes `issue_refund(order_id='ORD-502', method='original_payment', amount_usd=280.0, ...)`, and summarizes the approved refund.* + +- `Customer CUST-10 wants to return order ORD-101 (unopened USB-C cable delivered 5 days ago) for a refund.` + + *The agent looks up the order and customer profile and confirms the item is unopened within the 30-day return window. Because the default escalation policy asks the agent to consult before committing to a decision, the agent usually still calls `model_consult` once or twice here, the advisor confirms the straightforward decision, and `max_uses=2` caps the number of consultations in the turn. The agent then processes the full `$19.00` refund to `original_payment`.* + +## Graph + +```mermaid +graph TD + Agent[order_support_agent] -->|calls| GetOrder(get_order) + Agent -->|calls| GetProfile(get_customer_profile) + Agent -->|calls| Consult(model_consult / ModelConsultTool) + Agent -->|calls| IssueRefund(issue_refund) +``` + +## How To + +Define your domain tools, `get_order`, `get_customer_profile`, and `issue_refund`, and attach `ModelConsultTool` to the `Agent`: + +```python +from google.adk import Agent +from google.adk.tools import ModelConsultTool + + +def get_order(order_id: str) -> dict[str, str | int | float | bool | None]: + """Looks up an order by its identifier. + + Args: + order_id: Order identifier such as 'ORD-101' or 'ORD-502'. + + Returns: + A dictionary with the order details and any active defect bulletin. + """ + return { + "order_id": order_id, + "price_usd": 280.0, + "days_since_delivery": 45, + "opened": True, + "defect_bulletin": ( + "SB-2026-04: 90-day warranty replacement or store credit; cash refund" + " past 30 days requires Gold-tier loyalty exemption." + ), + } + + +def get_customer_profile(customer_id: str) -> dict[str, str | int | float]: + """Looks up a customer's loyalty tier and return history. + + Args: + customer_id: Customer identifier such as 'CUST-10' or 'CUST-108'. + + Returns: + A dictionary with the customer's loyalty tier and return rate percentage. + """ + return {"customer_id": customer_id, "tier": "gold", "return_rate_pct": 2.4} + + +def issue_refund( + order_id: str, + method: str, + amount_usd: float, + reason: str, +) -> dict[str, str | float]: + """Issues a refund or replacement for an order. + + Args: + order_id: Order identifier being refunded. + method: One of 'original_payment', 'store_credit', or 'replacement'. + amount_usd: Dollar amount to refund. + reason: Short explanation of the policy rule applied. + + Returns: + A confirmation record for the processed refund. + """ + return { + "status": "processed", + "order_id": order_id, + "method": method, + "amount_usd": amount_usd, + "reason": reason, + } + + +root_agent = Agent( + name="order_support_agent", + instruction=( + "You are an e-commerce order support assistant. Look up the order and" + " customer profile before calling issue_refund, and summarize the" + " outcome for the customer." + ), + tools=[ + get_order, + get_customer_profile, + issue_refund, + ModelConsultTool( + max_uses=2, + session_max_uses=5, + thinking_level="high", + ), + ], +) +``` + +Run the sample interactively from the repository root with the ADK CLI: + +```bash +adk run contributing/samples/tools/model_consult +``` + +Or launch the ADK web UI pointed at `contributing/samples/tools` and select `model_consult`: + +```bash +adk web contributing/samples/tools +``` + +## Related Guides + +- [ModelConsultTool and ModelConsultContextConfig](../../../../docs/guides/tools/model_consult/model_consult_tool/index.md) - Escalating hard decisions mid-generation to a stronger advisor model with per-turn and session budgets. diff --git a/contributing/samples/tools/model_consult/__init__.py b/contributing/samples/tools/model_consult/__init__.py new file mode 100644 index 00000000000..4015e47d6e4 --- /dev/null +++ b/contributing/samples/tools/model_consult/__init__.py @@ -0,0 +1,15 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from . import agent diff --git a/contributing/samples/tools/model_consult/agent.py b/contributing/samples/tools/model_consult/agent.py new file mode 100644 index 00000000000..a84ca779eb6 --- /dev/null +++ b/contributing/samples/tools/model_consult/agent.py @@ -0,0 +1,157 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Order support and refund policy sample using ModelConsultTool.""" + +from __future__ import annotations + +from google.adk import Agent +from google.adk.tools import ModelConsultTool + + +def get_order(order_id: str) -> dict[str, str | int | float | bool | None]: + """Looks up an order by its identifier. + + Args: + order_id: Order identifier such as 'ORD-101' or 'ORD-502'. + + Returns: + A dictionary with the order details and any active defect bulletin. + """ + orders = { + 'ORD-101': { + 'order_id': 'ORD-101', + 'customer_id': 'CUST-10', + 'item': 'USB-C Braided Cable', + 'category': 'accessories', + 'price_usd': 19.0, + 'days_since_delivery': 5, + 'opened': False, + 'defect_bulletin': None, + }, + 'ORD-502': { + 'order_id': 'ORD-502', + 'customer_id': 'CUST-108', + 'item': 'ProNC Wireless Headphones (Batch 2026-B)', + 'category': 'electronics', + 'price_usd': 280.0, + 'days_since_delivery': 45, + 'opened': True, + 'defect_bulletin': ( + 'SB-2026-04: Batch 2026-B battery drain defect — eligible for' + ' 90-day warranty replacement or full store credit; cash refund' + ' past 30 days requires Gold-tier loyalty exemption.' + ), + }, + } + return orders.get( + order_id, {'order_id': order_id, 'error': f'Order {order_id!r} not found'} + ) + + +def get_customer_profile(customer_id: str) -> dict[str, str | int | float]: + """Looks up a customer's loyalty tier and return history. + + Args: + customer_id: Customer identifier such as 'CUST-10' or 'CUST-108'. + + Returns: + A dictionary with the customer's loyalty tier and return rate percentage. + """ + customers = { + 'CUST-10': { + 'customer_id': 'CUST-10', + 'tier': 'standard', + 'lifetime_orders': 3, + 'return_rate_pct': 0.0, + }, + 'CUST-108': { + 'customer_id': 'CUST-108', + 'tier': 'gold', + 'lifetime_orders': 42, + 'return_rate_pct': 2.4, + }, + } + return customers.get( + customer_id, + { + 'customer_id': customer_id, + 'error': f'Customer {customer_id!r} not found', + }, + ) + + +def issue_refund( + order_id: str, + method: str, + amount_usd: float, + reason: str, +) -> dict[str, str | float]: + """Issues a refund or replacement for an order. + + Args: + order_id: Order identifier being refunded. + method: One of 'original_payment', 'store_credit', or 'replacement'. + amount_usd: Dollar amount to refund (use 0.0 for 'replacement'). + reason: Short explanation of the policy rule applied. + + Returns: + A confirmation record for the processed refund. + """ + return { + 'status': 'processed', + 'order_id': order_id, + 'method': method, + 'amount_usd': round(amount_usd, 2), + 'reason': reason, + } + + +_TASK_INSTRUCTION = """\ +You are an e-commerce order support assistant. Handle refund requests according +to the store's policy: +- Unopened items within 30 days of delivery qualify for a full + `original_payment` refund. +- Opened electronics within 30 days incur a 15% restocking fee (refund 85% of + `price_usd`), unless covered by an active `defect_bulletin`. +- Returns past 30 days are normally declined, with two exceptions: + 1. Items with an active `defect_bulletin` qualify for `replacement` or full + `store_credit` up to 90 days after delivery. + 2. `gold` tier customers with `return_rate_pct < 5.0` may convert a + defect-bulletin store credit into a full `original_payment` refund with no + restocking fee. + +Always call `get_order` and `get_customer_profile` to gather the order and +loyalty facts before calling `issue_refund`, and then summarize the outcome for +the customer. +""" + +root_agent = Agent( + name='order_support_agent', + description=( + 'Handles customer order returns, warranty defect bulletins, and loyalty' + ' refund policies.' + ), + instruction=_TASK_INSTRUCTION, + tools=[ + get_order, + get_customer_profile, + issue_refund, + ModelConsultTool( + max_uses=2, + session_max_uses=5, + thinking_level='high', + ), + ], +) diff --git a/docs/guides/README.md b/docs/guides/README.md index 2138378b741..56d76cb53df 100644 --- a/docs/guides/README.md +++ b/docs/guides/README.md @@ -116,6 +116,7 @@ This directory contains specific developer guides for the ADK Python implementat * [TelemetryConfig](telemetry/telemetry_config/index.md) - What ADK puts in its OpenTelemetry traces, and whether the text of prompts and replies is copied onto exported spans. ### Tools +* [ModelConsultTool and ModelConsultContextConfig](tools/model_consult/model_consult_tool/index.md) - Escalating hard decisions mid-generation to a stronger advisor model, with per-turn and session budgets. * [Node as tool](tools/node_tool/index.md) - Exposing workflows and deterministic nodes as agent tools with isolated runtime branching and resume support. * [to_mcp_server](tools/mcp_tool/agent_to_mcp/index.md) - Expose an ADK agent as an MCP server so any MCP host can drive it as a single tool (the MCP counterpart of to_a2a). diff --git a/docs/guides/tools/model_consult/model_consult_tool/index.md b/docs/guides/tools/model_consult/model_consult_tool/index.md new file mode 100644 index 00000000000..1728942f981 --- /dev/null +++ b/docs/guides/tools/model_consult/model_consult_tool/index.md @@ -0,0 +1,195 @@ +# ModelConsultTool + +`ModelConsultTool` gives a primary executor agent a callable tool named `model_consult` that escalates hard reasoning steps mid-generation to a stronger advisor model. The advisor model reviews the current session history, executor instructions, and available tool inventory with its own tool calling disabled, then returns structured guidance that the executor uses to continue the turn. + +## Introduction + +Many agent workloads consist mostly of routine steps such as reading files, querying logs, or formatting data, punctuated by one or two high-stakes decisions such as diagnosing a multi-service outage or reconciling multi-clause policy rules. Running every turn on a frontier reasoning model increases latency and token cost across the entire conversation, while running exclusively on a smaller model risks errors on harder reasoning steps. + +`ModelConsultTool` separates execution from deliberation inside a single agent turn. Your primary `Agent` runs on a fast model and handles tool execution and user responses directly. A default escalation policy tells the executor to call `model_consult` before committing to a decision, when stuck, and before declaring a task done, so even simple tasks usually trigger one consultation. `max_uses` and `session_max_uses` cap how often the executor can consult, and `executor_instruction` replaces the default policy with your own guidance. + +## Get started + +Attach `ModelConsultTool` to an `Agent` alongside your domain tools: + +```python +from google.adk import Agent +from google.adk.tools import ModelConsultTool + + +def lookup_order(order_id: str) -> dict[str, str]: + """Looks up order status by identifier.""" + return {"order_id": order_id, "status": "held_for_fraud_review"} + + +root_agent = Agent( + name="support_executor", + instruction=( + "You are an order support assistant. Resolve customer issues using" + " your tools." + ), + tools=[ + lookup_order, + ModelConsultTool( + max_uses=2, + session_max_uses=5, + thinking_level="high", + ), + ], +) +``` + +When `ModelConsultTool` prepares each outgoing executor request, it registers the `model_consult` function declaration and automatically appends a default escalation policy to the executor's system instruction so the executor knows when and how to consult the advisor. When `support_executor` invokes `model_consult(question="Should I release order ORD-42?")`, `ModelConsultTool` packages the session events, the executor's task instruction, and the names and descriptions of sibling tools such as `lookup_order` into a single advisor consultation. + +## How it works + +When the executor calls `model_consult`, `ModelConsultTool` performs four steps and returns a structured dictionary to the executor: + +1. **Budget verification** — `ModelConsultTool` checks the per-turn counter against `max_uses` and the session-wide counter against `session_max_uses`. If either cap has been reached, the tool returns `"status": "limit_reached"` immediately without calling the advisor model, and instructs the executor to proceed with the information already gathered. +1. **Context handover** — `ModelConsultTool` builds the advisor conversation from the non-partial, non-rewound events in `Session.events` according to `ModelConsultContextConfig`. Prior tool calls and tool responses in the session are flattened into readable text summaries so the advisor sees what actions have already been taken and what they returned, while any in-flight `model_consult` call is excluded. `ModelConsultTool` appends a final user handoff turn containing the active agent name, the executor's `question`, and any extra `context` string passed by the executor, and attaches the resolved executor instruction and sibling tool inventory to the advisor's system instruction when `include_agent_instruction` and `include_tool_inventory` are `True`. +1. **Tool-less advisor call** — `ModelConsultTool` calls the configured advisor `BaseLlm` with tool calling disabled and the default advisor system instruction, or a custom `advisor_instruction` when provided. Because tool declarations are excluded from the advisor request, the advisor cannot execute tools or produce side effects on its own; it can only return text guidance naming which tools the executor should invoke next and with what arguments. +1. **Structured tool response** — `ModelConsultTool` never raises an exception back into the agent loop: + - `"ok"`: Increments both usage counters and returns `"guidance"`, `"advisor_model"`, `"thinking_level"`, `"consults"` budget metadata, token `"usage"` counts, and `"latency_ms"`. + - `"limit_reached"`: Returned when `max_uses` or `session_max_uses` is already exhausted, with `"message"` and `"consults"`. + - `"error"`: Returned when the advisor call times out, fails, or produces no visible text, with `"error"`, `"message"`, `"advisor_model"`, and the current `"consults"` counters without incrementing them. + - `"invalid_request"`: Returned with `"message"` when `question` is empty or whitespace-only, without consuming budget. + +A successful consultation returns the following dictionary structure: + +```python +{ + "status": "ok", + "guidance": "1. Call lookup_order with order_id='ORD-42'.", + "advisor_model": "gemini-3.1-pro-preview", + "thinking_level": "high", + "consults": { + "used_this_turn": 1, + "max_uses": 2, + "used_this_session": 1, + "session_max_uses": 5, + "remaining": 1, + }, + "usage": { + "prompt_tokens": 612, + "output_tokens": 184, + "thoughts_tokens": 320, + "cached_tokens": 0, + "total_tokens": 1116, + }, + "latency_ms": 842.5, +} +``` + +## Configuration options + +`ModelConsultTool` configures advisor model selection, consultation budgets, and prompt overrides, while `ModelConsultContextConfig` controls how session events are formatted and bounded before handover. + +### ModelConsultTool options + +`ModelConsultTool` accepts the following constructor arguments: + +| Option | Type | Default | Description | +| :-------------------------- | :------------------------------------ | :------------------------- | :------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | +| `model` | `str \| BaseLlm` | `'gemini-3.1-pro-preview'` | Advisor model name resolved through ADK's model registry, or a pre-configured `BaseLlm` instance. | +| `max_uses` | `int \| None` | `None` | Maximum successful consultations per user turn. `None` means no per-turn cap. | +| `session_max_uses` | `int \| None` | `None` | Maximum successful consultations across the entire session. `None` means no session-wide cap. | +| `thinking_level` | `str \| types.ThinkingLevel \| None` | `'high'` | Reasoning effort for the advisor model: `'minimal'`, `'low'`, `'medium'`, `'high'`, a `types.ThinkingLevel` enum value, or `'off'`, `'none'`, or `None` to leave thinking unset. | +| `max_output_tokens` | `int \| None` | `None` | Optional cap on advisor output tokens, covering both visible output and thinking tokens on reasoning models. | +| `timeout_seconds` | `float \| None` | `None` | Per-call wall-clock timeout in seconds. `None` means no tool-level timeout. | +| `context_config` | `ModelConsultContextConfig \| None` | `None` | Controls how session history is packaged and bounded for the advisor. | +| `executor_instruction` | `str \| None` | `None` | Overrides the default escalation policy automatically appended to the executor's `system_instruction`. Pass `""` to disable automatic injection. | +| `advisor_instruction` | `str \| None` | `None` | Overrides the default system instruction sent to the advisor model. | +| `description` | `str \| None` | `None` | Overrides the default tool description shown to the executor model. | +| `include_agent_instruction` | `bool` | `True` | Forwards the executor agent's own instruction to the advisor so guidance respects the executor's constraints. | +| `include_tool_inventory` | `bool` | `True` | Includes the names and descriptions of the executor's other tools in the advisor system instruction. | +| `generate_content_config` | `types.GenerateContentConfig \| None` | `None` | Base generation config cloned per advisor call, such as `temperature` or `safety_settings`. | +| `name` | `str` | `'model_consult'` | Tool name exposed to the executor model. | + +`model` accepts either a model identifier string such as `'gemini-3.1-pro-preview'` or any `BaseLlm` instance, including `LiteLlm` wrappers for third-party models. + +`max_uses`, `session_max_uses`, `max_output_tokens`, and `timeout_seconds` enforce positive caps when set. Passing `0` or a negative number raises `ValueError` at construction time. Only successful advisor calls with `"status": "ok"` consume consultation budget; failed calls return `"status": "error"` without incrementing either counter. Call `has_remaining_budget(context)` with a `ToolContext` or `CallbackContext` to check whether at least one consultation remains in the current turn and session. + +`thinking_level` accepts `'minimal'`, `'low'`, `'medium'`, `'high'`, `'off'`, `'none'`, `''`, `None`, or a `types.ThinkingLevel` enum value. Passing `'off'`, `'none'`, `''`, or `None` leaves the advisor's thinking configuration unset. If the target advisor model rejects the thinking configuration as unsupported, `ModelConsultTool` automatically retries the call once without it. + +`max_output_tokens` caps the advisor's total generated tokens, including reasoning tokens on thinking models. If generation stops at `max_output_tokens` after producing partial text, `ModelConsultTool` appends a notice to the returned guidance; if thinking consumes the entire cap before any visible text is emitted, the call returns `"status": "error"`. Setting different values on `max_output_tokens` and `generate_content_config.max_output_tokens` raises `ValueError` at construction time. + +`executor_instruction`, `advisor_instruction`, and `description` override the built-in prompts that steer when the executor escalates and how the advisor formats its response. When `name` is customized without a custom `executor_instruction`, `ModelConsultTool` substitutes the custom tool name into the default escalation policy and scopes its per-turn and per-session state counters to `name`. + +`include_agent_instruction` and `include_tool_inventory` control whether the executor's resolved instruction and sibling tool list are appended to the advisor's system instruction. `generate_content_config` supplies a base `types.GenerateContentConfig` that is cloned for each advisor call with tool calling cleared. + +### ModelConsultContextConfig options + +`ModelConsultContextConfig` controls how `Session.events` is converted into the advisor's input contents: + +| Option | Type | Default | Description | +| :----------------- | :------------ | :--------- | :-------------------------------------------------------------------------------------------------------------------------------- | +| `mode` | `ContextMode` | `'events'` | `'events'` preserves multi-turn `types.Content` structure; `'transcript'` flattens history into a single text transcript. | +| `include_session` | `bool` | `True` | Sends the converted `Session.events` history when `True`, or only the `question` and `context` tool arguments when `False`. | +| `max_events` | `int \| None` | `None` | Keeps at most this many of the most recent non-partial session events before character budgeting. `None` keeps all events. | +| `max_chars` | `int \| None` | `200000` | Character budget across all handed-over session turns. `None` disables the character budget. | +| `max_part_chars` | `int` | `4000` | Per-part character cap on rendered tool calls, tool results, and code blocks, with plain text parts allowed eight times this cap. | +| `include_media` | `bool` | `True` | Forwards inline media and file references in `'events'` mode when `True`, or replaces them with text placeholders when `False`. | +| `include_thoughts` | `bool` | `False` | Includes the executor's internal thought parts in the advisor handover when `True`. | + +`ModelConsultContextConfig` validates fields strictly and rejects unknown keyword arguments or non-positive limits, requiring `max_events`, `max_chars`, and `max_part_chars` to be at least `1` when set. + +`mode` selects how session history is formatted for the advisor. `'events'` preserves alternating `user` and `model` `types.Content` turns, while `'transcript'` renders the history into a single labeled text block inside the user prompt for text-only or strict-alternation models. Setting `include_session=False` skips prior `Session.events` altogether so the advisor sees only the `question` and `context` tool arguments. + +`max_events` slices the most recent non-partial, non-rewound session events before part filtering and character budgeting. When the converted history exceeds `max_chars`, `ModelConsultTool` reserves up to one quarter of `max_chars` for leading turns so the initial goal remains visible when it fits, inserts a gap marker for dropped middle turns, and fills the remaining budget with the most recent turns. The newest turn is always kept and shortened in place if it exceeds the remaining character budget on its own. + +`max_part_chars` caps each rendered tool call argument string, tool response body, executable code snippet, and code execution result, while plain text and thought parts receive eight times `max_part_chars`. `include_media` forwards inline binary media and file references in `'events'` mode when `True`, or replaces them with text descriptors when `False`. `include_thoughts` defaults to `False` so the executor's internal reasoning does not anchor the advisor; when `True`, thought parts are prefixed with a thought marker. + +## Advanced applications + +The following patterns adapt `ModelConsultTool` for long-horizon sessions with large tool payloads or custom advisor model adapters. + +### Customizing context handover budgets + +For long-running debugging sessions with verbose tool outputs, pass a custom `ModelConsultContextConfig` to tighten per-part limits or switch to `'transcript'` mode for text-only advisor models: + +```python +from google.adk.tools import ModelConsultContextConfig +from google.adk.tools import ModelConsultTool + +consult_tool = ModelConsultTool( + max_uses=2, + session_max_uses=6, + context_config=ModelConsultContextConfig( + mode="transcript", + max_events=25, + max_chars=24000, + max_part_chars=3000, + include_media=False, + ), +) +``` + +### Supplying a custom BaseLlm advisor + +You can pass any `BaseLlm` instance to `ModelConsultTool(model=...)` when the advisor requires custom client options, Vertex AI credentials, or a non-Gemini model adapter: + +```python +from google.adk.models.google_llm import Gemini +from google.adk.tools import ModelConsultTool +from google.genai import types + +advisor_llm = Gemini(model="gemini-3.1-pro-preview") + +consult_tool = ModelConsultTool( + model=advisor_llm, + thinking_level="high", + max_output_tokens=4096, + generate_content_config=types.GenerateContentConfig( + temperature=0.2, + ), +) +``` + +## Limitations + +- **Advisory-only execution** — The advisor model runs with tool calling disabled and cannot invoke tools or mutate session state directly. The executor model must translate the advisor's guidance into concrete tool calls or user responses. +- **Shared token budget on reasoning models** — On Gemini reasoning models, `max_output_tokens` caps the sum of internal thinking tokens and visible output tokens. Setting `max_output_tokens` too low while `thinking_level='high'` can exhaust the token budget during thinking and return `"status": "error"` with zero visible guidance. Leave `max_output_tokens=None` or allocate sufficient headroom for both reasoning and output. + +## Related samples + +- [Model Consult Sample](../../../../../contributing/samples/tools/model_consult/agent.py) — E-commerce order support agent that combines `get_order`, `get_customer_profile`, and `issue_refund` with `ModelConsultTool` for multi-rule refund policy decisions. From 6f3003907564e93e625477f1856a483937aae3d6 Mon Sep 17 00:00:00 2001 From: Abhay Joshi Date: Tue, 29 Sep 2026 17:37:30 -0700 Subject: [PATCH 27/29] fix: isolate and clean up single_turn LlmAgent node_input events Merge https://github.com/google/adk-python/pull/7320 Fixes #7227 PiperOrigin-RevId: 990632092 --- .../adk/flows/llm_flows/context/_contents.py | 19 ++ src/google/adk/workflow/_llm_agent_wrapper.py | 25 +- .../workflow/test_llm_agent_as_node.py | 237 ++++++++++++++++++ 3 files changed, 274 insertions(+), 7 deletions(-) diff --git a/src/google/adk/flows/llm_flows/context/_contents.py b/src/google/adk/flows/llm_flows/context/_contents.py index cf6138ab581..f1d373ec2a4 100644 --- a/src/google/adk/flows/llm_flows/context/_contents.py +++ b/src/google/adk/flows/llm_flows/context/_contents.py @@ -126,6 +126,7 @@ async def run_async( agent.name, preserve_function_call_ids=preserve_function_call_ids, isolation_scope=invocation_context.isolation_scope, + node_path=invocation_context.node_path, is_single_turn=is_single_turn, user_content=invocation_context.user_content, include_thoughts_from_other_agents=include_thoughts_from_other_agents, @@ -139,6 +140,7 @@ async def run_async( agent.name, preserve_function_call_ids=preserve_function_call_ids, isolation_scope=invocation_context.isolation_scope, + node_path=invocation_context.node_path, is_single_turn=is_single_turn, user_content=invocation_context.user_content, include_thoughts_from_other_agents=False, @@ -312,6 +314,7 @@ def _should_include_event_in_context( event: Event, isolation_scope: str | None = None, *, + node_path: str | None = None, include_thoughts: bool = False, ) -> bool: """Determines if an event should be included in the LLM context. @@ -330,6 +333,7 @@ def _should_include_event_in_context( current_branch: The current branch of the agent. event: The event to filter. isolation_scope: The agent's isolation_scope. None means unscoped. + node_path: The current workflow node path, if executing as a node. Returns: True if the event should be included in the context, False otherwise. @@ -337,6 +341,15 @@ def _should_include_event_in_context( ev_iso = getattr(event, 'isolation_scope', None) if ev_iso != isolation_scope: return False + ev_node_info = getattr(event, 'node_info', None) + ev_node_path = getattr(ev_node_info, 'path', None) if ev_node_info else None + if ( + event.author == 'user' + and not event.get_function_responses() + and ev_node_path + and ev_node_path != (node_path or '') + ): + return False return not ( _contains_empty_content(event, include_thoughts=include_thoughts) or not _is_event_belongs_to_branch(current_branch, event) @@ -399,6 +412,7 @@ def _get_contents( *, preserve_function_call_ids: bool = False, isolation_scope: str | None = None, + node_path: str | None = None, is_single_turn: bool = False, user_content: types.Content | None = None, include_thoughts_from_other_agents: bool = False, @@ -414,6 +428,7 @@ def _get_contents( preserve_function_call_ids: Whether to preserve function call ids. isolation_scope: scope tag — when set, restricts events to those with matching ``event.isolation_scope`` (or unscoped). + node_path: The current workflow node path, if executing as a node. user_content: Fallback first user turn for task agents whose originating delegation FC is not in session (workflow-node task case). @@ -440,6 +455,7 @@ def _get_contents( current_branch, e, isolation_scope=isolation_scope, + node_path=node_path, include_thoughts=( include_thoughts_from_other_agents and _is_other_agent_reply(agent_name, e) @@ -586,6 +602,7 @@ def _get_current_turn_contents( preserve_function_call_ids: bool = False, is_single_turn: bool = False, isolation_scope: str | None = None, + node_path: str | None = None, user_content: types.Content | None = None, include_thoughts_from_other_agents: bool = False, ) -> list[types.Content]: @@ -637,6 +654,7 @@ def _get_current_turn_contents( current_branch, event, isolation_scope=isolation_scope, + node_path=node_path, include_thoughts=( include_thoughts_from_other_agents and _is_other_agent_reply(agent_name, event) @@ -652,6 +670,7 @@ def _get_current_turn_contents( agent_name, preserve_function_call_ids=preserve_function_call_ids, isolation_scope=isolation_scope, + node_path=node_path, is_single_turn=is_single_turn, user_content=user_content, include_thoughts_from_other_agents=include_thoughts_from_other_agents, diff --git a/src/google/adk/workflow/_llm_agent_wrapper.py b/src/google/adk/workflow/_llm_agent_wrapper.py index 4818b7a1ffa..13893e88b5d 100644 --- a/src/google/adk/workflow/_llm_agent_wrapper.py +++ b/src/google/adk/workflow/_llm_agent_wrapper.py @@ -310,7 +310,7 @@ def prepare_llm_agent_context(agent: LlmAgent, ctx: Context) -> Context: def prepare_llm_agent_input( agent: LlmAgent, ctx: Context, node_input: object -) -> None: +) -> Event | None: """Prepares the input for running LlmAgent as a node. For ``single_turn`` mode, append a user-role event with the input @@ -341,11 +341,14 @@ def prepare_llm_agent_input( or agent.mode != 'single_turn' or bool(ctx.resume_inputs) ): - return + return None agent_input = to_user_content(node_input) user_event = Event(author='user', message=agent_input) if user_event.content is not None: user_event.content.role = 'user' + node_path = getattr(ctx, 'node_path', None) + if isinstance(node_path, str) and node_path: + user_event.node_info.path = node_path iso = getattr(ctx, 'isolation_scope', None) if iso: user_event.isolation_scope = iso @@ -353,6 +356,7 @@ def prepare_llm_agent_input( if branch: user_event.branch = branch ctx.session.events.append(user_event) + return user_event def process_llm_agent_output( @@ -411,7 +415,7 @@ async def run_llm_agent_as_node( agent.include_contents = 'none' agent_ctx = prepare_llm_agent_context(agent, ctx) - prepare_llm_agent_input(agent, agent_ctx, node_input) + injected_input_event = prepare_llm_agent_input(agent, agent_ctx, node_input) ic = agent_ctx.get_invocation_context() update: dict[str, object] = {'agent': agent} @@ -435,10 +439,17 @@ async def run_llm_agent_as_node( if agent.mode == 'single_turn': # is_live is always False here (single_turn forces non-live). - async with aclosing(agent.run_async(ic)) as run_iter: - async for event in run_iter: - process_llm_agent_output(agent, ctx, event) - yield event + try: + async with aclosing(agent.run_async(ic)) as run_iter: + async for event in run_iter: + process_llm_agent_output(agent, ctx, event) + yield event + finally: + if ( + injected_input_event is not None + and injected_input_event in agent_ctx.session.events + ): + agent_ctx.session.events.remove(injected_input_event) return if agent.mode == 'chat': diff --git a/tests/unittests/workflow/test_llm_agent_as_node.py b/tests/unittests/workflow/test_llm_agent_as_node.py index 15f70215a32..9aa9a251aca 100644 --- a/tests/unittests/workflow/test_llm_agent_as_node.py +++ b/tests/unittests/workflow/test_llm_agent_as_node.py @@ -38,6 +38,7 @@ from google.adk.tools.function_tool import FunctionTool from google.adk.tools.long_running_tool import LongRunningFunctionTool from google.adk.workflow import _llm_agent_wrapper as agent_wrapper +from google.adk.workflow import node from google.adk.workflow import START from google.adk.workflow._llm_agent_wrapper import process_llm_agent_output from google.adk.workflow._workflow import Workflow @@ -1905,3 +1906,239 @@ def test_process_llm_agent_output_blank_schema_response_writes_no_state(): assert event.output is None assert ctx.actions.state_delta == {} + + +@pytest.mark.asyncio +async def test_single_turn_node_input_does_not_leak_across_sequential_tools( + request: pytest.FixtureRequest, +): + """Single-turn node_input must not leak into root agent across tool turns.""" + from . import testing_utils + + fake_pdf = b'%PDF-1.4-FAKE-BYTES' + worker_model = testing_utils.MockModel.create( + responses=['worker-summary-a', 'worker-summary-b'] + ) + worker = LlmAgent( + name='worker', + model=worker_model, + instruction='Summarize the attached document.', + mode='single_turn', + ) + + @node(name='run_worker', rerun_on_resume=True) + async def run_worker(ctx: Context, node_input: str) -> Any: + return await ctx.run_node( + worker, + node_input=types.Content( + role='user', + parts=[ + types.Part.from_text(text=f'INTERNAL-{node_input}'), + types.Part.from_bytes( + data=fake_pdf, mime_type='application/pdf' + ), + ], + ), + ) + + wf = Workflow( + name='doc_wf', + edges=[(START, run_worker)], + ) + + async def the_tool(label: str, tool_context: Context) -> dict[str, Any]: + out = await tool_context.run_node( + wf, node_input=label, run_id=f'run-{label}' + ) + assert not any( + ev.author == 'user' + and ev.content + and any( + p.text and 'INTERNAL-' in p.text for p in ev.content.parts or [] + ) + for ev in tool_context.session.events + ) + return {'summary': f'done:{label}:{out}'} + + fc_a = types.Part.from_function_call(name='the_tool', args={'label': 'a'}) + fc_b = types.Part.from_function_call(name='the_tool', args={'label': 'b'}) + root_model = testing_utils.MockModel.create( + responses=[fc_a, fc_b, 'All tools completed.'] + ) + root_agent = LlmAgent( + name='root_agent', + model=root_model, + instruction='Call the_tool twice sequentially.', + tools=[the_tool], + ) + + runner = _new_workflow_runner(root_agent, request.function.__name__) + await runner.run_async(testing_utils.get_user_content('run both tools')) + + # Worker received both inputs (text + inline PDF bytes). + assert len(worker_model.requests) == 2 + for expected_label, req in zip(['a', 'b'], worker_model.requests): + worker_texts = [ + p.text + for c in req.contents + for p in c.parts or [] + if p.text is not None + ] + worker_blobs = [ + p.inline_data.data + for c in req.contents + for p in c.parts or [] + if p.inline_data is not None + ] + assert any(f'INTERNAL-{expected_label}' in t for t in worker_texts) + assert fake_pdf in worker_blobs + + # Root agent made 3 LLM calls (initial -> after tool a -> after tool b). + # None of its requests should contain the worker's text or inline PDF. + assert len(root_model.requests) == 3 + for req in root_model.requests: + root_texts = [ + p.text + for c in req.contents + for p in c.parts or [] + if p.text is not None + ] + root_blobs = [ + p.inline_data + for c in req.contents + for p in c.parts or [] + if p.inline_data is not None + ] + assert not any('INTERNAL-' in t for t in root_texts) + assert not root_blobs + + +@pytest.mark.asyncio +async def test_parallel_single_turn_nodes_only_see_own_node_input( + request: pytest.FixtureRequest, +): + """Concurrent single_turn nodes sharing session.events see only own input.""" + import asyncio + + from . import testing_utils + + b_entered = asyncio.Event() + + async def wait_for_b(callback_context: Context) -> None: + del callback_context + await b_entered.wait() + + async def signal_b(callback_context: Context) -> None: + del callback_context + b_entered.set() + await asyncio.sleep(0) + + model_a = testing_utils.MockModel.create(responses=['out-a']) + model_b = testing_utils.MockModel.create(responses=['out-b']) + worker_a = LlmAgent( + name='worker_a', + model=model_a, + instruction='Worker A.', + mode='single_turn', + before_agent_callback=wait_for_b, + ) + worker_b = LlmAgent( + name='worker_b', + model=model_b, + instruction='Worker B.', + mode='single_turn', + before_agent_callback=signal_b, + ) + + @node(rerun_on_resume=True) + async def fanout(ctx: Context) -> dict[str, Any]: + res_a, res_b = await asyncio.gather( + ctx.run_node(worker_a, node_input='SECRET_FOR_A'), + ctx.run_node(worker_b, node_input='SECRET_FOR_B'), + ) + return {'a': res_a, 'b': res_b} + + wf = Workflow(name='parallel_wf', edges=[(START, fanout)]) + runner = _new_workflow_runner(wf, request.function.__name__) + await runner.run_async(testing_utils.get_user_content('start')) + + assert len(model_a.requests) == 1 + texts_a = [ + p.text + for c in model_a.requests[0].contents + for p in c.parts or [] + if p.text + ] + assert any('SECRET_FOR_A' in t for t in texts_a) + assert not any('SECRET_FOR_B' in t for t in texts_a) + + assert len(model_b.requests) == 1 + texts_b = [ + p.text + for c in model_b.requests[0].contents + for p in c.parts or [] + if p.text + ] + assert any('SECRET_FOR_B' in t for t in texts_b) + assert not any('SECRET_FOR_A' in t for t in texts_b) + + +@pytest.mark.asyncio +async def test_synthesized_task_fr_preserved_across_node_paths( + request: pytest.FixtureRequest, +): + """Synthesized task FunctionResponse (author='user') survives node_path changes.""" + from . import testing_utils + + task_fc = types.Part.from_function_call( + name='specialist', + args={'request': 'compute'}, + ) + finish_fc = types.Part.from_function_call( + name='finish_task', + args={'result': '42'}, + ) + specialist_model = testing_utils.MockModel.create(responses=[finish_fc]) + specialist = LlmAgent( + name='specialist', + model=specialist_model, + instruction='Specialist.', + mode='task', + ) + coord_model = testing_utils.MockModel.create( + responses=[task_fc, 'First pass done.', 'Second pass done.'] + ) + coordinator = LlmAgent( + name='coordinator', + model=coord_model, + instruction='Coordinator.', + mode='chat', + sub_agents=[specialist], + ) + + @node(rerun_on_resume=True) + async def step_one(ctx: Context) -> Any: + return await ctx.run_node(coordinator) + + @node(rerun_on_resume=True) + async def step_two(ctx: Context) -> Any: + return await ctx.run_node(coordinator) + + wf = Workflow( + name='loop_wf', + edges=[(START, step_one), (step_one, step_two)], + ) + runner = _new_workflow_runner(wf, request.function.__name__) + await runner.run_async(testing_utils.get_user_content('go')) + + assert len(coord_model.requests) == 3 + last_req = coord_model.requests[-1] + frs = [ + p.function_response + for c in last_req.contents + for p in c.parts or [] + if p.function_response is not None + ] + assert len(frs) == 1 + assert frs[0].name == 'specialist' + assert frs[0].response == {'result': '42'} From d312c0eca3707dd162b21bbe00a2325758ea3d6a Mon Sep 17 00:00:00 2001 From: Jason Zhang Date: Tue, 29 Sep 2026 20:19:00 -0700 Subject: [PATCH 28/29] refactor(tools): extract shared URL validation helpers into _url_validator.py Extract the SSRF and URL target validation helpers (`_parse_request_target`, `_is_blocked_hostname`, `_is_blocked_address`, `_embedded_ipv4`, `_resolve_direct_addresses`, and `_reject_blocked_proxied_hostname`) into `_url_validator.py`. Previously, `ComputerUseToolset._wrap_navigate_with_url_validation` lazily imported private helpers from `load_web_page` inside the wrapper function to avoid pulling `requests` into the computer-use import path. Moving the validation logic into `_url_validator.py` allows `load_web_page`, `ComputerUseToolset`, and other outbound HTTP tools to import the shared validation helpers directly at the module level without extra runtime dependencies. Test files that patched `load_web_page.socket` were updated after private helpers were pruned from the `load_web_page` namespace. Co-authored-by: Jason Zhang PiperOrigin-RevId: 990692502 --- src/google/adk/tools/_url_validator.py | 222 +++++++++++++++++ .../computer_use/computer_use_toolset.py | 8 +- src/google/adk/tools/load_web_page.py | 212 +---------------- .../computer_use/test_computer_use_toolset.py | 4 +- tests/unittests/tools/test_load_web_page.py | 19 +- tests/unittests/tools/test_url_validator.py | 225 ++++++++++++++++++ 6 files changed, 472 insertions(+), 218 deletions(-) create mode 100644 src/google/adk/tools/_url_validator.py create mode 100644 tests/unittests/tools/test_url_validator.py diff --git a/src/google/adk/tools/_url_validator.py b/src/google/adk/tools/_url_validator.py new file mode 100644 index 00000000000..4daa9c1833d --- /dev/null +++ b/src/google/adk/tools/_url_validator.py @@ -0,0 +1,222 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Url checks shared by the tools that open a url the model supplied.""" + +from __future__ import annotations + +from dataclasses import dataclass +import ipaddress +import socket +from urllib.parse import ParseResult +from urllib.parse import urlparse + +_ALLOWED_URL_SCHEMES = frozenset({'http', 'https'}) +_DEFAULT_PORT_BY_SCHEME = {'http': 80, 'https': 443} +# Hostnames that always designate the local machine or a metadata endpoint. +_BLOCKED_HOSTNAMES = frozenset({ + 'localhost', + 'metadata', + 'metadata.goog', +}) +# Hostname suffixes reserved for loopback, link-local and internal networks. +_BLOCKED_HOSTNAME_SUFFIXES = ( + '.localhost', + '.local', + '.internal', + '.metadata.goog', +) +_ResolvedAddress = ipaddress.IPv4Address | ipaddress.IPv6Address + + +@dataclass(frozen=True) +class _RequestTarget: + parsed_url: ParseResult + scheme: str + hostname: str + host_header: str + + +def _format_host(hostname: str) -> str: + if ':' in hostname: + return f'[{hostname}]' + return hostname + + +def _default_port_for_scheme(scheme: str) -> int: + return _DEFAULT_PORT_BY_SCHEME[scheme] + + +def _build_host_header( + *, hostname: str, scheme: str, explicit_port: int | None +) -> str: + formatted_hostname = _format_host(hostname) + if explicit_port is None or explicit_port == _default_port_for_scheme(scheme): + return formatted_hostname + return f'{formatted_hostname}:{explicit_port}' + + +def _parse_request_target(url: str) -> _RequestTarget: + parsed_url = urlparse(url) + scheme = parsed_url.scheme.lower() + if scheme not in _ALLOWED_URL_SCHEMES: + raise ValueError(f'Unsupported url scheme: {url}') + + hostname = parsed_url.hostname + if not hostname: + raise ValueError(f'URL is missing a hostname: {url}') + + try: + explicit_port = parsed_url.port + except ValueError as exc: + raise ValueError(f'Invalid url port: {url}') from exc + + return _RequestTarget( + parsed_url=parsed_url, + scheme=scheme, + hostname=hostname, + host_header=_build_host_header( + hostname=hostname, + scheme=scheme, + explicit_port=explicit_port, + ), + ) + + +def _parse_ip_literal(hostname: str) -> _ResolvedAddress | None: + try: + return ipaddress.ip_address(hostname) + except ValueError: + return None + + +def _is_blocked_hostname(hostname: str) -> bool: + """Reports whether a name designates loopback or internal infrastructure. + + This check is purely lexical, so unlike the address checks it also applies + when an outbound proxy performs the DNS resolution on our behalf. + + Args: + hostname: The hostname parsed out of the requested url. + + Returns: + True if the request must be refused without contacting the host. + """ + normalized_hostname = hostname.rstrip('.').lower() + if normalized_hostname in _BLOCKED_HOSTNAMES: + return True + return normalized_hostname.endswith(_BLOCKED_HOSTNAME_SUFFIXES) + + +_NAT64_WELL_KNOWN_PREFIX = ipaddress.ip_network('64:ff9b::/96') + + +def _embedded_ipv4(address: _ResolvedAddress) -> ipaddress.IPv4Address | None: + """Returns the IPv4 address embedded in an IPv6 address, if any. + + ``is_global`` on the outer IPv6 address does not reflect the reachability of + the embedded IPv4 target for IPv4-mapped (``::ffff:a.b.c.d``), IPv4-compatible + (``::a.b.c.d``), 6to4 (``2002::/16``) and NAT64 (``64:ff9b::/96``) addresses. + For example ``64:ff9b::169.254.169.254`` is reported as global but, on a + network with NAT64, routes to the internal ``169.254.169.254`` metadata + endpoint. Returning the embedded IPv4 lets the caller vet it directly. + """ + if not isinstance(address, ipaddress.IPv6Address): + return None + if address.ipv4_mapped is not None: + return address.ipv4_mapped + if address.sixtofour is not None: + return address.sixtofour + if address in _NAT64_WELL_KNOWN_PREFIX: + return ipaddress.IPv4Address(int(address) & 0xFFFFFFFF) + # IPv4-compatible ``::a.b.c.d`` (deprecated): top 96 bits zero, low 32 bits a + # non-trivial IPv4 (excluding ``::`` and ``::1``). + packed = int(address) + if packed >> 32 == 0 and (packed & 0xFFFFFFFF) not in (0, 1): + return ipaddress.IPv4Address(packed & 0xFFFFFFFF) + return None + + +def _is_blocked_address(address: _ResolvedAddress) -> bool: + if not address.is_global: + return True + # Reject IPv6 addresses that embed a non-global IPv4 target (NAT64, + # IPv4-compatible, etc.), which `is_global` alone does not catch. + embedded = _embedded_ipv4(address) + return embedded is not None and not embedded.is_global + + +def _resolve_host_addresses(hostname: str) -> tuple[_ResolvedAddress, ...]: + resolved_address = _parse_ip_literal(hostname) + + if resolved_address is not None: + return (resolved_address,) + + try: + address_info = socket.getaddrinfo( + hostname, + None, + type=socket.SOCK_STREAM, + proto=socket.IPPROTO_TCP, + ) + except (socket.gaierror, UnicodeError) as exc: + raise ValueError(f'Unable to resolve host: {hostname}') from exc + + resolved_addresses: list[_ResolvedAddress] = [] + for family, _, _, _, sockaddr in address_info: + if family not in (socket.AF_INET, socket.AF_INET6): + continue + resolved_addresses.append(ipaddress.ip_address(sockaddr[0])) + + if not resolved_addresses: + raise ValueError(f'Unable to resolve host: {hostname}') + + return tuple(resolved_addresses) + + +def _resolve_direct_addresses(hostname: str) -> tuple[_ResolvedAddress, ...]: + resolved_addresses = tuple(dict.fromkeys(_resolve_host_addresses(hostname))) + if any(_is_blocked_address(address) for address in resolved_addresses): + raise ValueError(f'Blocked host: {hostname}') + return resolved_addresses + + +def _reject_blocked_proxied_hostname(hostname: str) -> None: + """Best-effort address check for a hostname that the proxy will resolve. + + The proxy performs the authoritative DNS resolution and opens the connection, + so the local lookup here is advisory rather than a pin. It still refuses the + common case where a public resolver maps the requested name onto a metadata, + loopback or otherwise private address. + + A local resolution failure is not treated as an error: split-horizon DNS and + egress-only networks legitimately leave the proxy as the only resolver, and + failing closed there would break every such deployment. Those environments + are covered by `_is_blocked_hostname` instead. A proxy that resolves a + public-looking name to an internal address remains outside what a client can + detect, and has to be constrained by the proxy's own egress policy. + + Args: + hostname: The hostname that will be handed to the proxy. + + Raises: + ValueError: If the local resolver maps the hostname to a non-global + address. + """ + try: + resolved_addresses = _resolve_host_addresses(hostname) + except ValueError: + return + if any(_is_blocked_address(address) for address in resolved_addresses): + raise ValueError(f'Blocked host: {hostname}') diff --git a/src/google/adk/tools/computer_use/computer_use_toolset.py b/src/google/adk/tools/computer_use/computer_use_toolset.py index 103ba73ea2f..f9a579a580e 100644 --- a/src/google/adk/tools/computer_use/computer_use_toolset.py +++ b/src/google/adk/tools/computer_use/computer_use_toolset.py @@ -31,6 +31,9 @@ from ...features import experimental from ...features import FeatureName from ...models.llm_request import LlmRequest +from .._url_validator import _is_blocked_hostname +from .._url_validator import _parse_request_target +from .._url_validator import _resolve_direct_addresses from ..base_toolset import BaseToolset from ..tool_context import ToolContext from .base_computer import BaseComputer @@ -135,11 +138,6 @@ def _wrap_navigate_with_url_validation( @functools.wraps(navigate_method) async def wrapper(url: str) -> Any: - # Deferred to keep `requests` off the computer-use import path. - from ..load_web_page import _is_blocked_hostname - from ..load_web_page import _parse_request_target - from ..load_web_page import _resolve_direct_addresses - try: if not isinstance(url, str): raise ValueError("url is not a string") diff --git a/src/google/adk/tools/load_web_page.py b/src/google/adk/tools/load_web_page.py index 0da5f0ccdd9..ccfbad565eb 100644 --- a/src/google/adk/tools/load_web_page.py +++ b/src/google/adk/tools/load_web_page.py @@ -16,21 +16,25 @@ """Tool for web browse.""" -from dataclasses import dataclass -import ipaddress -import socket import time from typing import Any from urllib.parse import ParseResult -from urllib.parse import urlparse import requests from requests.adapters import HTTPAdapter from requests.utils import get_environ_proxies from requests.utils import select_proxy -_ALLOWED_URL_SCHEMES = frozenset({'http', 'https'}) -_DEFAULT_PORT_BY_SCHEME = {'http': 80, 'https': 443} +from ._url_validator import _format_host +from ._url_validator import _is_blocked_address +from ._url_validator import _is_blocked_hostname +from ._url_validator import _parse_ip_literal +from ._url_validator import _parse_request_target +from ._url_validator import _reject_blocked_proxied_hostname +from ._url_validator import _RequestTarget +from ._url_validator import _resolve_direct_addresses +from ._url_validator import _ResolvedAddress + # Default timeout in seconds for HTTP requests. This bounds the connect phase # and the gap between two received chunks, but not the total transfer time. _DEFAULT_TIMEOUT_SECONDS = 30 @@ -41,28 +45,6 @@ _MAX_RESPONSE_BYTES = 10 * 1024 * 1024 # Chunk size used while streaming a response body. _RESPONSE_CHUNK_BYTES = 64 * 1024 -# Hostnames that always designate the local machine or a metadata endpoint. -_BLOCKED_HOSTNAMES = frozenset({ - 'localhost', - 'metadata', - 'metadata.goog', -}) -# Hostname suffixes reserved for loopback, link-local and internal networks. -_BLOCKED_HOSTNAME_SUFFIXES = ( - '.localhost', - '.local', - '.internal', - '.metadata.goog', -) -_ResolvedAddress = ipaddress.IPv4Address | ipaddress.IPv6Address - - -@dataclass(frozen=True) -class _RequestTarget: - parsed_url: ParseResult - scheme: str - hostname: str - host_header: str class _PinnedAddressAdapter(HTTPAdapter): @@ -120,185 +102,11 @@ def _failed_to_fetch_message(url: str) -> str: return f'Failed to fetch url: {url}' -def _format_host(hostname: str) -> str: - if ':' in hostname: - return f'[{hostname}]' - return hostname - - -def _default_port_for_scheme(scheme: str) -> int: - return _DEFAULT_PORT_BY_SCHEME[scheme] - - -def _build_host_header( - *, hostname: str, scheme: str, explicit_port: int | None -) -> str: - formatted_hostname = _format_host(hostname) - if explicit_port is None or explicit_port == _default_port_for_scheme(scheme): - return formatted_hostname - return f'{formatted_hostname}:{explicit_port}' - - -def _parse_request_target(url: str) -> _RequestTarget: - parsed_url = urlparse(url) - scheme = parsed_url.scheme.lower() - if scheme not in _ALLOWED_URL_SCHEMES: - raise ValueError(f'Unsupported url scheme: {url}') - - hostname = parsed_url.hostname - if not hostname: - raise ValueError(f'URL is missing a hostname: {url}') - - try: - explicit_port = parsed_url.port - except ValueError as exc: - raise ValueError(f'Invalid url port: {url}') from exc - - return _RequestTarget( - parsed_url=parsed_url, - scheme=scheme, - hostname=hostname, - host_header=_build_host_header( - hostname=hostname, - scheme=scheme, - explicit_port=explicit_port, - ), - ) - - -def _parse_ip_literal(hostname: str) -> _ResolvedAddress | None: - try: - return ipaddress.ip_address(hostname) - except ValueError: - return None - - -def _is_blocked_hostname(hostname: str) -> bool: - """Reports whether a name designates loopback or internal infrastructure. - - This check is purely lexical, so unlike the address checks it also applies - when an outbound proxy performs the DNS resolution on our behalf. - - Args: - hostname: The hostname parsed out of the requested url. - - Returns: - True if the request must be refused without contacting the host. - """ - normalized_hostname = hostname.rstrip('.').lower() - if normalized_hostname in _BLOCKED_HOSTNAMES: - return True - return normalized_hostname.endswith(_BLOCKED_HOSTNAME_SUFFIXES) - - -_NAT64_WELL_KNOWN_PREFIX = ipaddress.ip_network('64:ff9b::/96') - - -def _embedded_ipv4(address: _ResolvedAddress) -> ipaddress.IPv4Address | None: - """Returns the IPv4 address embedded in an IPv6 address, if any. - - ``is_global`` on the outer IPv6 address does not reflect the reachability of - the embedded IPv4 target for IPv4-mapped (``::ffff:a.b.c.d``), IPv4-compatible - (``::a.b.c.d``), 6to4 (``2002::/16``) and NAT64 (``64:ff9b::/96``) addresses. - For example ``64:ff9b::169.254.169.254`` is reported as global but, on a - network with NAT64, routes to the internal ``169.254.169.254`` metadata - endpoint. Returning the embedded IPv4 lets the caller vet it directly. - """ - if not isinstance(address, ipaddress.IPv6Address): - return None - if address.ipv4_mapped is not None: - return address.ipv4_mapped - if address.sixtofour is not None: - return address.sixtofour - if address in _NAT64_WELL_KNOWN_PREFIX: - return ipaddress.IPv4Address(int(address) & 0xFFFFFFFF) - # IPv4-compatible ``::a.b.c.d`` (deprecated): top 96 bits zero, low 32 bits a - # non-trivial IPv4 (excluding ``::`` and ``::1``). - packed = int(address) - if packed >> 32 == 0 and (packed & 0xFFFFFFFF) not in (0, 1): - return ipaddress.IPv4Address(packed & 0xFFFFFFFF) - return None - - -def _is_blocked_address(address: _ResolvedAddress) -> bool: - if not address.is_global: - return True - # Reject IPv6 addresses that embed a non-global IPv4 target (NAT64, - # IPv4-compatible, etc.), which `is_global` alone does not catch. - embedded = _embedded_ipv4(address) - return embedded is not None and not embedded.is_global - - -def _resolve_host_addresses(hostname: str) -> tuple[_ResolvedAddress, ...]: - resolved_address = _parse_ip_literal(hostname) - - if resolved_address is not None: - return (resolved_address,) - - try: - address_info = socket.getaddrinfo( - hostname, - None, - type=socket.SOCK_STREAM, - proto=socket.IPPROTO_TCP, - ) - except (socket.gaierror, UnicodeError) as exc: - raise ValueError(f'Unable to resolve host: {hostname}') from exc - - resolved_addresses: list[_ResolvedAddress] = [] - for family, _, _, _, sockaddr in address_info: - if family not in (socket.AF_INET, socket.AF_INET6): - continue - resolved_addresses.append(ipaddress.ip_address(sockaddr[0])) - - if not resolved_addresses: - raise ValueError(f'Unable to resolve host: {hostname}') - - return tuple(resolved_addresses) - - def _get_proxy_url(url: str) -> str | None: proxies = get_environ_proxies(url) return select_proxy(url, proxies) -def _resolve_direct_addresses(hostname: str) -> tuple[_ResolvedAddress, ...]: - resolved_addresses = tuple(dict.fromkeys(_resolve_host_addresses(hostname))) - if any(_is_blocked_address(address) for address in resolved_addresses): - raise ValueError(f'Blocked host: {hostname}') - return resolved_addresses - - -def _reject_blocked_proxied_hostname(hostname: str) -> None: - """Best-effort address check for a hostname that the proxy will resolve. - - The proxy performs the authoritative DNS resolution and opens the connection, - so the local lookup here is advisory rather than a pin. It still refuses the - common case where a public resolver maps the requested name onto a metadata, - loopback or otherwise private address. - - A local resolution failure is not treated as an error: split-horizon DNS and - egress-only networks legitimately leave the proxy as the only resolver, and - failing closed there would break every such deployment. Those environments - are covered by `_is_blocked_hostname` instead. A proxy that resolves a - public-looking name to an internal address remains outside what a client can - detect, and has to be constrained by the proxy's own egress policy. - - Args: - hostname: The hostname that will be handed to the proxy. - - Raises: - ValueError: If the local resolver maps the hostname to a non-global - address. - """ - try: - resolved_addresses = _resolve_host_addresses(hostname) - except ValueError: - return - if any(_is_blocked_address(address) for address in resolved_addresses): - raise ValueError(f'Blocked host: {hostname}') - - def _declared_content_length(response: requests.Response) -> int: """Returns the declared body size, or 0 when the header is unusable. diff --git a/tests/unittests/tools/computer_use/test_computer_use_toolset.py b/tests/unittests/tools/computer_use/test_computer_use_toolset.py index 48278e8ae7c..bb59c59897c 100644 --- a/tests/unittests/tools/computer_use/test_computer_use_toolset.py +++ b/tests/unittests/tools/computer_use/test_computer_use_toolset.py @@ -18,7 +18,7 @@ from unittest.mock import Mock from google.adk.models.llm_request import LlmRequest -from google.adk.tools import load_web_page +from google.adk.tools import _url_validator # Use the actual ComputerEnvironment enum from the code from google.adk.tools.computer_use.base_computer import BaseComputer from google.adk.tools.computer_use.base_computer import ComputerEnvironment @@ -641,7 +641,7 @@ def resolver(self, monkeypatch) -> Mock: ("93.184.216.34", 0), )] ) - monkeypatch.setattr(load_web_page.socket, "getaddrinfo", resolver) + monkeypatch.setattr(_url_validator.socket, "getaddrinfo", resolver) return resolver @staticmethod diff --git a/tests/unittests/tools/test_load_web_page.py b/tests/unittests/tools/test_load_web_page.py index de907fb9700..3711c2a2bd9 100644 --- a/tests/unittests/tools/test_load_web_page.py +++ b/tests/unittests/tools/test_load_web_page.py @@ -20,6 +20,7 @@ from unittest import mock from google.adk.tools import load_web_page as load_web_page_module +import google.adk.tools._url_validator as url_validator_module import pytest import requests @@ -59,7 +60,7 @@ def _set_proxy_env(monkeypatch): def _mock_getaddrinfo(monkeypatch, *addresses: str): monkeypatch.setattr( - load_web_page_module.socket, + url_validator_module.socket, 'getaddrinfo', mock.Mock( return_value=[ @@ -214,7 +215,7 @@ def _send( def test_load_web_page_blocks_private_hostname_targets(monkeypatch): _clear_proxy_env(monkeypatch) monkeypatch.setattr( - load_web_page_module.socket, + url_validator_module.socket, 'getaddrinfo', mock.Mock( return_value=[( @@ -247,7 +248,7 @@ def test_load_web_page_uses_proxy_for_unresolved_public_hostnames(monkeypatch): # Split-horizon DNS and egress-only networks leave the proxy as the only # resolver, so a local lookup failure must not block the request. monkeypatch.setattr( - load_web_page_module.socket, + url_validator_module.socket, 'getaddrinfo', mock.Mock(side_effect=socket.gaierror('no such host')), ) @@ -326,7 +327,7 @@ def test_load_web_page_blocks_internal_hostnames_behind_a_proxy( """Internal names are rejected lexically, without relying on local DNS.""" _set_proxy_env(monkeypatch) monkeypatch.setattr( - load_web_page_module.socket, + url_validator_module.socket, 'getaddrinfo', mock.Mock(side_effect=AssertionError('unexpected local DNS lookup')), ) @@ -359,7 +360,7 @@ def test_load_web_page_fetches_public_urls_by_pinning_the_resolved_ip( ): _clear_proxy_env(monkeypatch) monkeypatch.setattr( - load_web_page_module.socket, + url_validator_module.socket, 'getaddrinfo', mock.Mock( return_value=[( @@ -412,7 +413,7 @@ def test_load_web_page_tries_another_resolved_address_after_connect_error( ): _clear_proxy_env(monkeypatch) monkeypatch.setattr( - load_web_page_module.socket, + url_validator_module.socket, 'getaddrinfo', mock.Mock( return_value=[ @@ -480,7 +481,7 @@ def test_load_web_page_passes_timeout_to_pinned_session(monkeypatch): """Verify that the default timeout is passed to the pinned IP session.""" _clear_proxy_env(monkeypatch) monkeypatch.setattr( - load_web_page_module.socket, + url_validator_module.socket, 'getaddrinfo', mock.Mock( return_value=[( @@ -530,7 +531,7 @@ def test_load_web_page_passes_timeout_to_proxied_get(monkeypatch): """Verify that the default timeout is passed to requests.get when proxy is used.""" _set_proxy_env(monkeypatch) monkeypatch.setattr( - load_web_page_module.socket, + url_validator_module.socket, 'getaddrinfo', mock.Mock(side_effect=socket.gaierror('no such host')), ) @@ -556,7 +557,7 @@ def test_load_web_page_returns_failure_on_timeout(monkeypatch): """Verify that a timeout exception is converted to a failed to fetch message.""" _clear_proxy_env(monkeypatch) monkeypatch.setattr( - load_web_page_module.socket, + url_validator_module.socket, 'getaddrinfo', mock.Mock( return_value=[( diff --git a/tests/unittests/tools/test_url_validator.py b/tests/unittests/tools/test_url_validator.py new file mode 100644 index 00000000000..edaa2f5c734 --- /dev/null +++ b/tests/unittests/tools/test_url_validator.py @@ -0,0 +1,225 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from __future__ import annotations + +import ipaddress +import socket +from unittest import mock + +from google.adk.tools._url_validator import _embedded_ipv4 +from google.adk.tools._url_validator import _is_blocked_address +from google.adk.tools._url_validator import _is_blocked_hostname +from google.adk.tools._url_validator import _parse_request_target +from google.adk.tools._url_validator import _resolve_direct_addresses +from google.adk.tools._url_validator import _resolve_host_addresses +import google.adk.tools._url_validator as url_validator +import pytest + +_PUBLIC_IPV4 = '8.8.8.8' +_PUBLIC_IPV6 = '2001:4860:4860::8888' + +_LOOPBACK_IPV4 = '127.0.0.1' +_LOOPBACK_IPV6 = '::1' +_PRIVATE_IPV4 = '10.0.0.1' +_METADATA_IPV4 = '169.254.169.254' + +# IPv6 addresses embedding an IPv4. `ipaddress.is_global` does not always +# account for the embedded address, so the validator must check it. +_PUBLIC_VIA_NAT64 = f'64:ff9b::{_PUBLIC_IPV4}' +_METADATA_VIA_NAT64 = f'64:ff9b::{_METADATA_IPV4}' +_LOOPBACK_VIA_IPV4_MAPPED = f'::ffff:{_LOOPBACK_IPV4}' +_METADATA_VIA_IPV4_COMPATIBLE = f'::{_METADATA_IPV4}' +_METADATA_VIA_6TO4 = '2002:a9fe:a9fe::' # 169.254.169.254 in hex. + + +def _fake_dns(monkeypatch: pytest.MonkeyPatch, *ipv4s: str) -> None: + """Makes every DNS lookup return the given IPv4 addresses.""" + records = [ + (socket.AF_INET, socket.SOCK_STREAM, socket.IPPROTO_TCP, '', (ip, 0)) + for ip in ipv4s + ] + monkeypatch.setattr( + url_validator.socket, 'getaddrinfo', mock.Mock(return_value=records) + ) + + +def _broken_dns(monkeypatch: pytest.MonkeyPatch, error: Exception) -> None: + """Makes every DNS lookup raise `error`.""" + monkeypatch.setattr( + url_validator.socket, 'getaddrinfo', mock.Mock(side_effect=error) + ) + + +# --- _parse_request_target --------------------------------------------------- + + +@pytest.mark.parametrize( + ('url', 'expected_hostname', 'expected_host_header'), + [ + ('https://example.com:443/path', 'example.com', 'example.com'), + ('http://example.com:8080/path', 'example.com', 'example.com:8080'), + ( + f'http://[{_PUBLIC_IPV6}]:8080/', + _PUBLIC_IPV6, + f'[{_PUBLIC_IPV6}]:8080', + ), + ], +) +def test_parse_request_target_accepts_http_urls( + url: str, expected_hostname: str, expected_host_header: str +): + target = _parse_request_target(url) + + assert target.hostname == expected_hostname + assert target.host_header == expected_host_header + + +@pytest.mark.parametrize( + ('url', 'expected_error'), + [ + ('file:///etc/passwd', 'Unsupported url scheme'), + ('http:///missing-host', 'missing a hostname'), + ('http://example.com:99999/', 'Invalid url port'), + ], +) +def test_parse_request_target_rejects_invalid_urls( + url: str, expected_error: str +): + with pytest.raises(ValueError, match=expected_error): + _parse_request_target(url) + + +# --- _is_blocked_hostname ---------------------------------------------------- + + +@pytest.mark.parametrize( + 'hostname', + [ + 'localhost', + 'LOCALHOST.', + 'a.localhost', + 'metadata', + 'metadata.goog', + 'sub.metadata.goog', + 'instance.internal', + 'service.local', + ], +) +def test_is_blocked_hostname_blocks_internal_names(hostname: str): + assert _is_blocked_hostname(hostname) + + +@pytest.mark.parametrize('hostname', ['example.com', 'localhost.example.com']) +def test_is_blocked_hostname_allows_other_names(hostname: str): + assert not _is_blocked_hostname(hostname) + + +# --- _embedded_ipv4 ---------------------------------------------------------- + + +@pytest.mark.parametrize( + ('ip', 'expected'), + [ + (_LOOPBACK_VIA_IPV4_MAPPED, _LOOPBACK_IPV4), + (_METADATA_VIA_6TO4, _METADATA_IPV4), + (_METADATA_VIA_NAT64, _METADATA_IPV4), + (_METADATA_VIA_IPV4_COMPATIBLE, _METADATA_IPV4), + ], +) +def test_embedded_ipv4_extracts_the_wrapped_address(ip: str, expected: str): + assert _embedded_ipv4(ipaddress.ip_address(ip)) == ipaddress.ip_address( + expected + ) + + +@pytest.mark.parametrize( + 'ip', [_PUBLIC_IPV4, _PUBLIC_IPV6, '::', _LOOPBACK_IPV6] +) +def test_embedded_ipv4_returns_none_without_an_embedded_address(ip: str): + assert _embedded_ipv4(ipaddress.ip_address(ip)) is None + + +# --- _is_blocked_address ----------------------------------------------------- + + +@pytest.mark.parametrize('ip', [_PUBLIC_IPV4, _PUBLIC_IPV6, _PUBLIC_VIA_NAT64]) +def test_is_blocked_address_allows_public_addresses(ip: str): + assert not _is_blocked_address(ipaddress.ip_address(ip)) + + +@pytest.mark.parametrize( + 'ip', [_LOOPBACK_IPV4, _LOOPBACK_IPV6, _PRIVATE_IPV4, _METADATA_IPV4] +) +def test_is_blocked_address_blocks_non_public_addresses(ip: str): + assert _is_blocked_address(ipaddress.ip_address(ip)) + + +@pytest.mark.parametrize( + 'ip', + [ + _METADATA_VIA_NAT64, + _LOOPBACK_VIA_IPV4_MAPPED, + _METADATA_VIA_IPV4_COMPATIBLE, + _METADATA_VIA_6TO4, + ], +) +def test_is_blocked_address_blocks_ipv6_wrapping_non_public_ipv4(ip: str): + assert _is_blocked_address(ipaddress.ip_address(ip)) + + +# --- _resolve_host_addresses ------------------------------------------------- + + +@pytest.mark.parametrize('ip', [_PUBLIC_IPV4, _PUBLIC_IPV6]) +def test_resolve_host_addresses_returns_ip_literal_without_dns( + monkeypatch, ip: str +): + _broken_dns(monkeypatch, AssertionError('unexpected DNS lookup')) + + assert _resolve_host_addresses(ip) == (ipaddress.ip_address(ip),) + + +def test_resolve_host_addresses_reports_dns_failure(monkeypatch): + _broken_dns(monkeypatch, socket.gaierror('Name or service not known')) + + with pytest.raises(ValueError, match='Unable to resolve host'): + _resolve_host_addresses('example.com') + + +# --- _resolve_direct_addresses ----------------------------------------------- + + +def test_resolve_direct_addresses_returns_unique_public_addresses( + monkeypatch, +): + _fake_dns(monkeypatch, _PUBLIC_IPV4, _PUBLIC_IPV4) + + assert _resolve_direct_addresses('example.com') == ( + ipaddress.ip_address(_PUBLIC_IPV4), + ) + + +def test_resolve_direct_addresses_blocks_host_with_any_non_public_address( + monkeypatch, +): + _fake_dns(monkeypatch, _PUBLIC_IPV4, _METADATA_IPV4) + + with pytest.raises(ValueError, match='Blocked host'): + _resolve_direct_addresses('example.com') + + +def test_resolve_direct_addresses_blocks_non_public_ip_literal(): + with pytest.raises(ValueError, match='Blocked host'): + _resolve_direct_addresses(_LOOPBACK_IPV4) From b8f50f8a2d5bad082624f7aa1965ac300a0afeab Mon Sep 17 00:00:00 2001 From: Haiyuan Cao Date: Tue, 29 Sep 2026 21:08:56 -0700 Subject: [PATCH 29/29] fix(plugins): chain re-raised interrupts to nothing, look up handle per call Review of the previous commit found two small hardening gaps: - The KeyboardInterrupt or SystemExit set aside while a content_formatter failure was described is raised after the row with `from None`. That only sets __suppress_context__: Python still records the exception the caller is handling, such as the error ADK passes to an error callback, as its __context__, so code that walks __context__ could reach it and its text. It is now raised while a context-free stand-in is handled, so that stand-in is its only link. - The stand-in wrapper bound Logger.handle when the plugin was imported, so a later class-level patch of Logger.handle, as instrumentation and test fixtures apply, never saw this logger's records. The class's handle is now looked up on each call, still inside the stand-in. The content_formatter docstring also states the downstream effect plainly: the row's status is left as the event set it, usually 'OK', so a query that counts any non-NULL error_message as an error, such as the BigQuery Agent Analytics SDK's error predicate, counts a formatter failure as an error. Behavior is unchanged. Tests: the re-raised interrupt links to neither the caller's exception nor anything chained to it (KeyboardInterrupt and SystemExit), and a class-level patch of Logger.handle made after import sees the plugin's formatter warning, handled while the stand-in is. Both fail before this change. Refs: https://github.com/GoogleCloudPlatform/BigQuery-Agent-Analytics-SDK/issues/485 (item A) Co-Authored-By: Claude Opus 5.5 (1M context) --- .../bigquery_agent_analytics_plugin.py | 31 ++++-- .../test_bigquery_agent_analytics_plugin.py | 95 +++++++++++++++++++ 2 files changed, 117 insertions(+), 9 deletions(-) diff --git a/src/google/adk/plugins/bigquery_agent_analytics_plugin.py b/src/google/adk/plugins/bigquery_agent_analytics_plugin.py index bd9fd4b62aa..3d8e3926008 100644 --- a/src/google/adk/plugins/bigquery_agent_analytics_plugin.py +++ b/src/google/adk/plugins/bigquery_agent_analytics_plugin.py @@ -119,7 +119,9 @@ class _LoggingStandIn(Exception): logging's ``handleError``, and handlers that report the current exception, print the exception being handled. Without this stand-in that can be one the caller is handling, such as the error ADK passes to an error callback, - and its text can carry the content this plugin keeps out of logs. + and its text can carry the content this plugin keeps out of logs. For the + same reason, a set-aside interrupt is raised while a stand-in is handled, + so that the stand-in, not the caller's exception, is its ``__context__``. """ @@ -130,19 +132,19 @@ def _handle_records_with_a_stand_in(target: logging.Logger) -> None: while a record is handled can reach the caller's exception, even by a handler that ignores ``__suppress_context__``. Records are unchanged: ``Logger._log`` resolves the calling function and any ``exc_info`` before - it calls ``handle``. + it calls ``handle``. The class's ``handle`` is looked up on every call, so + a patch applied to it after import still reaches ``target``. Args: target: The logger whose records to handle this way. """ - handle = type(target).handle def handle_with_stand_in(record: logging.LogRecord) -> None: try: raise _LoggingStandIn except _LoggingStandIn as stand_in: stand_in.__context__ = None - handle(target, record) + type(target).handle(target, record) target.handle = handle_with_stand_in # type: ignore[method-assign] @@ -2426,7 +2428,10 @@ class BigQueryLoggerConfig: propagates, and no row is written. An event that already carries an ``error_message``, such as a ``TOOL_ERROR``, keeps it first, followed by ``; `` and the formatter - failure. + failure. The row's ``status`` is left as the event set it, usually + ``'OK'``, so a query that counts any non-NULL ``error_message`` as an + error, such as the BigQuery Agent Analytics SDK's error predicate, + counts a formatter failure as an error. gcs_bucket_name: GCS bucket for offloading large content. connection_id: BigQuery connection ID for ObjectRef columns. log_session_metadata: Whether to log session metadata. @@ -7476,8 +7481,9 @@ async def _log_event( Raises: KeyboardInterrupt: A signal handler, or a log handler or filter, raised one while a content_formatter failure was being described. - A new one without text is raised after the row was handed to the - writer; see ``_settle_formatter_outcome``. + A new one without text, chained to no exception the caller is + handling, is raised after the row was handed to the writer; see + ``_settle_formatter_outcome``. SystemExit: Likewise; it keeps the exit code only if that is an int. """ interrupts: list[BaseException] = [] @@ -7493,8 +7499,15 @@ async def _log_event( finally: if interrupts: # Raised only now, after the row was handed to the writer, so that - # neither the row nor the signal is lost. - raise interrupts[0] from None + # neither the row nor the signal is lost. Raising it while a + # context-free stand-in is handled makes the stand-in its + # __context__; `from None` alone would only hide the caller's + # exception from printers that honor __suppress_context__. + try: + raise _LoggingStandIn + except _LoggingStandIn as stand_in: + stand_in.__context__ = None + raise interrupts[0] from None async def _log_event_row( self, diff --git a/tests/unittests/plugins/test_bigquery_agent_analytics_plugin.py b/tests/unittests/plugins/test_bigquery_agent_analytics_plugin.py index 96a18180489..4f9e8fcfb4e 100644 --- a/tests/unittests/plugins/test_bigquery_agent_analytics_plugin.py +++ b/tests/unittests/plugins/test_bigquery_agent_analytics_plugin.py @@ -5017,6 +5017,101 @@ async def test_plugin_error_logs_keep_their_own_exception( assert len(records) == 1 assert records[0].exc_info[0] is RuntimeError + async def test_later_patches_of_logger_handle_see_plugin_records( + self, mock_write_client, invocation_context, dummy_arrow_schema + ): + """A patch of Logger.handle made after import applies to this logger too. + + Instrumentation and test fixtures patch the class. The patched handle + must still run while the stand-in is handled. + """ + plugin_module = bigquery_agent_analytics_plugin + plugin_logger = logging.getLogger("google_adk." + plugin_module.__name__) + original_handle = logging.Logger.handle + handled_while = [] + + def recording_handle(target, record): + if target is plugin_logger and record.getMessage().startswith( + "Content formatter " + ): + handled_while.append(sys.exc_info()[0]) + return original_handle(target, record) + + config = plugin_module.BigQueryLoggerConfig( + content_formatter=_raise_import_error + ) + with mock.patch.object(logging.Logger, "handle", recording_handle): + row, _ = await self._log_user_message_contained( + config, mock_write_client, invocation_context, dummy_arrow_schema + ) + + assert row["content"] == plugin_module._FORMATTER_FAILED_SENTINEL + assert handled_while == [plugin_module._LoggingStandIn] + + @pytest.mark.parametrize( + "interrupt", + [KeyboardInterrupt, SystemExit], + ids=["keyboard_interrupt", "system_exit"], + ) + async def test_raised_interrupt_is_not_chained_to_the_callers_exception( + self, + interrupt, + mock_write_client, + invocation_context, + dummy_arrow_schema, + ): + """The interrupt raised after the row links to nothing the caller handles. + + ``from None`` only hides a link from printers that honor + ``__suppress_context__``; code that walks ``__context__`` would still + reach the caller's exception and its text. + """ + plugin_logger = logging.getLogger( + "google_adk." + bigquery_agent_analytics_plugin.__name__ + ) + + class _InterruptingHandler(logging.Handler): + + def emit(self, record): + if record.getMessage().startswith("Content formatter "): + raise interrupt() + + config = bigquery_agent_analytics_plugin.BigQueryLoggerConfig( + content_formatter=_raise_import_error + ) + caller_failure = ValueError(f"the caller is handling {self.SECRET}") + handler = _InterruptingHandler() + plugin_logger.addHandler(handler) + try: + try: + raise caller_failure + except ValueError: + row, _, escaped = await self._log_user_message_catching( + config, mock_write_client, invocation_context, dummy_arrow_schema + ) + finally: + plugin_logger.removeHandler(handler) + + assert ( + row["content"] + == bigquery_agent_analytics_plugin._FORMATTER_FAILED_SENTINEL + ) + assert type(escaped) is interrupt + assert escaped.__suppress_context__ + chain = [] + link = escaped + while link is not None and len(chain) < 10: + chain.append(link) + link = link.__context__ + assert not any( + link is caller_failure for link in chain + ), "the interrupt is chained to the caller's exception" + # Only a constant stand-in, itself chained to nothing, may be linked. + assert [type(link) for link in chain[1:]] in ( + [], + [bigquery_agent_analytics_plugin._LoggingStandIn], + ) + @pytest.mark.parametrize( "rethrow", [True, False],