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 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/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/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. 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/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/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/src/google/adk/plugins/bigquery_agent_analytics_plugin.py b/src/google/adk/plugins/bigquery_agent_analytics_plugin.py index a4c6d71c589..0ee1b6d60e3 100644 --- a/src/google/adk/plugins/bigquery_agent_analytics_plugin.py +++ b/src/google/adk/plugins/bigquery_agent_analytics_plugin.py @@ -52,6 +52,8 @@ import sys import threading import time +from types import CoroutineType +from types import GeneratorType from types import MappingProxyType from types import TracebackType from typing import Any @@ -89,6 +91,7 @@ from google.genai import types from opentelemetry import trace import pydantic +from pydantic import BaseModel try: import pyarrow as pa @@ -113,7 +116,45 @@ 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. 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__``. + """ + + +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``. 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. + """ + + def handle_with_stand_in(record: logging.LogRecord) -> None: + try: + raise _LoggingStandIn + except _LoggingStandIn as stand_in: + stand_in.__context__ = None + type(target).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. @@ -1266,6 +1307,265 @@ 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, 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. + + Args: + cls: The class to label. + + Returns: + The label, or None when ``cls`` has none. + """ + 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. + + 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. + 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 + 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 ``, + with ```` in place of the label when the class cannot be + read at all. + """ + outcome = "raised" if raised else "returned unsupported type" + 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 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 ``_settle_formatter_outcome``, so the warning is + still logged. + + 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 BaseException: + return "[traceback could not be rendered]" + + +# SystemExit's own ``code`` descriptor; reading it through the descriptor +# runs no code of a SystemExit subclass. +_SYSTEM_EXIT_CODE = SystemExit.__dict__["code"] + + +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: + """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, +) -> tuple[Any, Optional[str], Optional[BaseException]]: + """Decides what the row logs after the content_formatter call, fail closed. + + 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; + - rendering the debug traceback; + - emitting the warning through whatever filters and handlers are + configured. + + 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 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 + 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. + 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: + ``(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} " + try: + if failure is not None: + failed_type: type = type(failure) + else: + 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, 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. + 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) + 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 # cannot catch graphs that create new objects per access (Mock-like duck # typing); the cap turns unbounded recursion into a redacted leaf. @@ -2509,7 +2809,43 @@ 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 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``. 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 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. 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. @@ -2579,6 +2915,23 @@ class BigQueryLoggerConfig: or if it cannot take those keyword arguments. One that raises an ``Exception`` or returns anything else is skipped at run time: the built-in rule applies, and the logged warning names only the tool. + 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. 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 @@ -2666,6 +3019,10 @@ class BigQueryLoggerConfig: credentials_identifier: Optional[str] = None # Application rules for tool results that report a failure without raising. tool_result_classifier: Optional[ToolResultClassifier] = None + # 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 # ============================================================================== @@ -4684,7 +5041,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( @@ -7593,6 +7951,57 @@ 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, 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] = [] + 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. 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, + 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 @@ -7657,55 +8066,52 @@ async def _log_event( is_truncated = True timestamp = datetime.now(timezone.utc) + formatter_error: Optional[str] = None if self.config.content_formatter: + formatted: Any = None + failure: Optional[Exception] = None try: formatted = self.config.content_formatter(raw_content, event_type) - if isinstance(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. - type(formatted) 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. The - # message is CONSTANT: even a class NAME can be payload-derived - # via type(name, ...). - logger.warning( - "Content formatter returned an unsupported result type for" - " event %s; writing sentinel instead of original content.", - event_type, - ) - 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, ...). - logger.warning( - "Content formatter failed for event %s; writing sentinel" - " instead of original content.", - event_type, - ) - raw_content = _FORMATTER_FAILED_SENTINEL + # payload. The exception message and traceback can embed the + # protected content, and even the class NAME can be payload-derived + # 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 + # Everything after the call runs behind _settle_formatter_outcome's + # 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, + debug=self.config.debug_content_formatter_errors, + ) + 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 + + # 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. 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 = ( + 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 @@ -7822,7 +8228,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/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/_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/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/src/google/adk/tools/model_consult/__init__.py b/src/google/adk/tools/model_consult/__init__.py new file mode 100644 index 00000000000..272f37623b9 --- /dev/null +++ b/src/google/adk/tools/model_consult/__init__.py @@ -0,0 +1,35 @@ +# 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 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/_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/src/google/adk/tools/model_consult/_context.py b/src/google/adk/tools/model_consult/_context.py new file mode 100644 index 00000000000..c110161a8b1 --- /dev/null +++ b/src/google/adk/tools/model_consult/_context.py @@ -0,0 +1,487 @@ +# 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 Literal +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 + +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. +_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', 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. + + 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 and config.mode != 'transcript': + return part + 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: + 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/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/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/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( 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/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", [ diff --git a/tests/unittests/plugins/test_bigquery_agent_analytics_plugin.py b/tests/unittests/plugins/test_bigquery_agent_analytics_plugin.py index 15686fdbea7..9db3f9e489a 100644 --- a/tests/unittests/plugins/test_bigquery_agent_analytics_plugin.py +++ b/tests/unittests/plugins/test_bigquery_agent_analytics_plugin.py @@ -19,10 +19,12 @@ import dataclasses import functools import gc +import io import json import logging import os import pickle +import signal import sys import threading import time @@ -3818,6 +3820,1652 @@ 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. +# 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): + """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__ + + +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() + + +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 _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 _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})") + + +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 _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") + + +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 + + 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_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]" + + return RedactedPayload() + + +@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 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." + ) + 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): + yield + globals().pop(self.PAYLOAD_IDENTIFIER, None) + + 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() + + 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 + ) + + 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 [ + record + for record in caplog.records + if record.getMessage().startswith("Content formatter ") + ] + + @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"), + [ + pytest.param( + _raise_import_error, + "content_formatter raised ImportError", + id="builtin_exception", + ), + pytest.param( + _raise_module_level_exception, + "content_formatter raised ", + 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( + _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", + id="unsupported_builtin_result", + ), + pytest.param( + _return_generator, + "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="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_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" + " ", + id="pydantic_model_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, + caplog, + ): + """A failed formatter's row fails closed and names a trusted class. + + 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 + ) + + 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 ( + row["content"] + == 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 + ): + """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( + ("error_callback", "error_text", "expected_error_message"), + [ + pytest.param( + "on_tool_error_callback", + "upstream timed out after 30s", + "upstream timed out after 30s;" + " content_formatter raised ImportError", + id="appended_after_existing_message", + ), + pytest.param( + # A model error is recorded as str(error), so an exception + # without a message leaves the event's own message empty. + "on_model_error_callback", + "", + "content_formatter raised ImportError", + id="empty_existing_message", + ), + pytest.param( + # A tool error without a message is recorded by its type name, + # which is the event's own message and so comes first. + "on_tool_error_callback", + "", + "RuntimeError; content_formatter raised ImportError", + id="message_less_tool_error", + ), + ], + ) + async def test_formatter_failure_follows_the_events_own_error_message( + self, + error_callback, + 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) + if error_callback == "on_tool_error_callback": + await plugin.on_tool_error_callback( + tool=tool, + tool_args={"account": self.SECRET}, + tool_context=tool_context, + error=RuntimeError(error_text), + ) + else: + await plugin.on_model_error_callback( + callback_context=tool_context, + llm_request=llm_request_lib.LlmRequest(model="gemini-pro"), + error=RuntimeError(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 is constant, 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 [record.getMessage() for record in warnings] == [ + self.DEFAULT_WARNING + ] + 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 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, + 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 + 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, KeyboardInterrupt, id="debug_on_interrupt"), + pytest.param(True, SystemExit, id="debug_on_exit"), + 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. + + 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) + + 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() + + 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 " + ) + 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( + "interrupt", + [KeyboardInterrupt, SystemExit, asyncio.CancelledError], + ids=["keyboard_interrupt", "system_exit", "cancelled_error"], + ) + async def test_interrupts_raised_by_the_formatter_call_still_propagate( + self, interrupt, mock_write_client, invocation_context, dummy_arrow_schema + ): + """Containment covers diagnosis only, never the formatter call itself. + + 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. + """ + + def formatter(content, event_type): + raise interrupt("raised by the formatter call") + + config = bigquery_agent_analytics_plugin.BigQueryLoggerConfig( + content_formatter=formatter, debug_content_formatter_errors=True + ) + + with pytest.raises(interrupt, match="raised by the formatter call"): + await self._log_user_message( + config, mock_write_client, invocation_context, dummy_arrow_schema + ) + + async def test_failing_log_handler_never_prints_the_formatter_exception( + self, + capsys, + monkeypatch, + mock_write_client, + invocation_context, + dummy_arrow_schema, + ): + """A handler that fails while warning cannot print the formatter error. + + 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) + config = bigquery_agent_analytics_plugin.BigQueryLoggerConfig( + content_formatter=_raise_import_error + ) + closed_stream = io.StringIO() + closed_stream.close() + + 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}" + ) + 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( + "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 + + 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], + 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", + [ + 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"), + 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( + self, + step, + raised, + injected, + mock_write_client, + invocation_context, + dummy_arrow_schema, + caplog, + ): + """Whatever any step after the formatter call raises, the row survives. + + 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. + + 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__) + + 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) + + 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=formatter, 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(), + "admit": mock.patch.object( + plugin_module, "_natively_parsed", raise_injected, create=True + ), + "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): + 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"): + 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: + 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: """Tests that _safe_callback prevents plugin errors from propagating.""" 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/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.""" 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_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() 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}) 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 + ) 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) 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'}