Files
adk-python/tests/unittests/plugins/test_bigquery_agent_analytics_plugin.py
T
Haiyuan CaoandCopybara-Service ea034877ec fix: Improve BigQuery Agent Analytics plugin reliability and code quality
This CL fixes several bugs in the BigQuery Agent Analytics plugin and refactors the internal data-passing pattern for better type safety and maintainability.

-   **Stale Loop State Validation:** Use `loop.is_closed()` — a public, reliable API — to detect and clean up stale asyncio loop states in `_batch_processor_prop`, `_get_loop_state`, and `flush`. The previous approach used `asyncio.Queue._loop` which is `None` on Python 3.10+, causing the check to always treat states as stale.
-   **Quota Project ID Fallback:** Remove the `or project_id` fallback when setting `quota_project_id` on `BigQueryWriteAsyncClient`. This fixes Workload Identity Federation flows where the federated identity lacks `serviceusage.services.use` on the quota project.
-   **Kwargs Passthrough:** Pass `**kwargs` through to `_log_event` in all callbacks. Previously only model callbacks forwarded them, causing custom attributes (e.g. `customer_id`) to silently drop for agent, tool, run, and error events.
-   **State Delta Logging:** Replace the dead `on_state_change_callback` (never invoked by the framework) with `on_event_callback`, which is already dispatched by the runner for every event. Remove duplicate `STATE_DELTA` logging from `after_tool_callback`.
-   **EventData Dataclass:** Replace the `**kwargs`-as-data-bus pattern in `_log_event` with an explicit `EventData` dataclass. This makes the interface self-documenting, catches typos at construction time, and eliminates shared dict mutation across `_resolve_span_ids`, `_extract_latency`, and `_enrich_attributes`. All 12 callback call sites now construct typed `EventData` instances.
-   **Multi-Subagent Tool Logging Tests:** Add `TestMultiSubagentToolLogging` (6 tests) verifying that tool events are correctly attributed to subagents in multi-turn, multi-agent scenarios. Total tests: 111 (up from 60).

Co-authored-by: Haiyuan Cao <haiyuan@google.com>
PiperOrigin-RevId: 871381533
2026-02-17 10:24:34 -08:00

3952 lines
135 KiB
Python

# 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 asyncio
import contextlib
import dataclasses
import json
from unittest import mock
from google.adk.agents import base_agent
from google.adk.agents import callback_context as callback_context_lib
from google.adk.agents import invocation_context as invocation_context_lib
from google.adk.events import event as event_lib
from google.adk.events import event_actions as event_actions_lib
from google.adk.models import llm_request as llm_request_lib
from google.adk.models import llm_response as llm_response_lib
from google.adk.plugins import bigquery_agent_analytics_plugin
from google.adk.plugins import plugin_manager as plugin_manager_lib
from google.adk.sessions import base_session_service as base_session_service_lib
from google.adk.sessions import session as session_lib
from google.adk.tools import base_tool as base_tool_lib
from google.adk.tools import tool_context as tool_context_lib
from google.adk.version import __version__
import google.auth
from google.auth import exceptions as auth_exceptions
import google.auth.credentials
from google.cloud import bigquery
from google.cloud import exceptions as cloud_exceptions
from google.genai import types
from opentelemetry import trace
import pyarrow as pa
import pytest
PROJECT_ID = "test-gcp-project"
DATASET_ID = "adk_logs"
TABLE_ID = "agent_events"
DEFAULT_STREAM_NAME = (
f"projects/{PROJECT_ID}/datasets/{DATASET_ID}/tables/{TABLE_ID}/_default"
)
# --- Pytest Fixtures ---
@pytest.fixture
def mock_session():
mock_s = mock.create_autospec(
session_lib.Session, instance=True, spec_set=True
)
type(mock_s).id = mock.PropertyMock(return_value="session-123")
type(mock_s).user_id = mock.PropertyMock(return_value="user-456")
type(mock_s).app_name = mock.PropertyMock(return_value="test_app")
type(mock_s).state = mock.PropertyMock(return_value={})
return mock_s
@pytest.fixture
def mock_agent():
mock_a = mock.create_autospec(
base_agent.BaseAgent, instance=True, spec_set=True
)
# Mock the 'name' property
type(mock_a).name = mock.PropertyMock(return_value="MyTestAgent")
type(mock_a).instruction = mock.PropertyMock(return_value="Test Instruction")
return mock_a
@pytest.fixture
def invocation_context(mock_agent, mock_session):
mock_session_service = mock.create_autospec(
base_session_service_lib.BaseSessionService, instance=True, spec_set=True
)
mock_plugin_manager = mock.create_autospec(
plugin_manager_lib.PluginManager, instance=True, spec_set=True
)
return invocation_context_lib.InvocationContext(
agent=mock_agent,
session=mock_session,
invocation_id="inv-789",
session_service=mock_session_service,
plugin_manager=mock_plugin_manager,
)
@pytest.fixture
def callback_context(invocation_context):
return callback_context_lib.CallbackContext(
invocation_context=invocation_context
)
@pytest.fixture
def tool_context(invocation_context):
return tool_context_lib.ToolContext(invocation_context=invocation_context)
@pytest.fixture
def mock_auth_default():
mock_creds = mock.create_autospec(
google.auth.credentials.Credentials, instance=True, spec_set=True
)
with mock.patch.object(
google.auth,
"default",
autospec=True,
return_value=(mock_creds, PROJECT_ID),
) as mock_auth:
yield mock_auth
@pytest.fixture
def mock_bq_client():
with mock.patch.object(bigquery, "Client", autospec=True) as mock_cls:
yield mock_cls.return_value
@pytest.fixture
def mock_write_client():
with mock.patch.object(
bigquery_agent_analytics_plugin, "BigQueryWriteAsyncClient", autospec=True
) as mock_cls:
mock_client = mock_cls.return_value
mock_client.transport = mock.AsyncMock()
async def fake_append_rows(requests, **kwargs):
# This function is now async, so `await client.append_rows` works.
mock_append_rows_response = mock.MagicMock()
mock_append_rows_response.row_errors = []
mock_append_rows_response.error = mock.MagicMock()
mock_append_rows_response.error.code = 0 # OK status
# This a gen is what's returned *after* the await.
return _async_gen(mock_append_rows_response)
mock_client.append_rows.side_effect = fake_append_rows
yield mock_client
@pytest.fixture
def dummy_arrow_schema():
return pa.schema([
pa.field("timestamp", pa.timestamp("us", tz="UTC"), nullable=False),
pa.field("root_agent_name", pa.string(), nullable=True),
pa.field("event_type", pa.string(), nullable=True),
pa.field("agent", pa.string(), nullable=True),
pa.field("session_id", pa.string(), nullable=True),
pa.field("invocation_id", pa.string(), nullable=True),
pa.field("user_id", pa.string(), nullable=True),
pa.field("trace_id", pa.string(), nullable=True),
pa.field("span_id", pa.string(), nullable=True),
pa.field("parent_span_id", pa.string(), nullable=True),
pa.field(
"content", pa.string(), nullable=True
), # JSON stored as string in Arrow
pa.field(
"content_parts",
pa.list_(
pa.struct([
pa.field("mime_type", pa.string(), nullable=True),
pa.field("uri", pa.string(), nullable=True),
pa.field(
"object_ref",
pa.struct([
pa.field("uri", pa.string(), nullable=True),
pa.field("authorizer", pa.string(), nullable=True),
pa.field("version", pa.string(), nullable=True),
pa.field(
"details",
pa.string(),
nullable=True,
metadata={
b"ARROW:extension:name": (
b"google:sqlType:json"
)
},
),
]),
nullable=True,
),
pa.field("text", pa.string(), nullable=True),
pa.field("part_index", pa.int64(), nullable=True),
pa.field("part_attributes", pa.string(), nullable=True),
pa.field("storage_mode", pa.string(), nullable=True),
])
),
nullable=True,
),
pa.field("attributes", pa.string(), nullable=True),
pa.field("latency_ms", pa.string(), nullable=True),
pa.field("status", pa.string(), nullable=True),
pa.field("error_message", pa.string(), nullable=True),
pa.field("is_truncated", pa.bool_(), nullable=True),
])
@pytest.fixture
def mock_to_arrow_schema(dummy_arrow_schema):
with mock.patch.object(
bigquery_agent_analytics_plugin,
"to_arrow_schema",
autospec=True,
return_value=dummy_arrow_schema,
) as mock_func:
yield mock_func
@pytest.fixture
def mock_asyncio_to_thread():
async def fake_to_thread(func, *args, **kwargs):
return func(*args, **kwargs)
with mock.patch(
"asyncio.to_thread", side_effect=fake_to_thread
) as mock_async:
yield mock_async
@pytest.fixture
def mock_storage_client():
with mock.patch("google.cloud.storage.Client") as mock_client:
yield mock_client
@pytest.fixture
async def bq_plugin_inst(
mock_auth_default,
mock_bq_client,
mock_write_client,
mock_to_arrow_schema,
mock_asyncio_to_thread,
):
plugin = bigquery_agent_analytics_plugin.BigQueryAgentAnalyticsPlugin(
project_id=PROJECT_ID,
dataset_id=DATASET_ID,
table_id=TABLE_ID,
)
await plugin._ensure_started() # Ensure clients are initialized
mock_write_client.append_rows.reset_mock()
yield plugin
await plugin.shutdown()
@contextlib.asynccontextmanager
async def managed_plugin(*args, **kwargs):
"""Async context manager to ensure plugin shutdown."""
plugin = bigquery_agent_analytics_plugin.BigQueryAgentAnalyticsPlugin(
*args, **kwargs
)
try:
yield plugin
finally:
await plugin.shutdown()
# --- Helper Functions ---
async def _async_gen(val):
yield val
async def _get_captured_event_dict_async(mock_write_client, expected_schema):
"""Helper to get the event_dict passed to append_rows."""
mock_write_client.append_rows.assert_called_once()
call_args = mock_write_client.append_rows.call_args
requests_iter = call_args.args[0]
requests = []
if hasattr(requests_iter, "__aiter__"):
async for req in requests_iter:
requests.append(req)
else:
requests = list(requests_iter)
assert len(requests) == 1
request = requests[0]
assert request.write_stream == DEFAULT_STREAM_NAME
assert request.trace_id == f"google-adk-bq-logger/{__version__}"
# Parse the Arrow batch back to a dict for verification
try:
reader = pa.ipc.open_stream(request.arrow_rows.rows.serialized_record_batch)
table = reader.read_all()
except Exception:
# Fallback: try reading as a single batch
buf = pa.py_buffer(request.arrow_rows.rows.serialized_record_batch)
batch = pa.ipc.read_record_batch(buf, expected_schema)
table = pa.Table.from_batches([batch])
assert table.schema.equals(
expected_schema
), f"Schema mismatch: Expected {expected_schema}, got {table.schema}"
pydict = table.to_pydict()
return {k: v[0] for k, v in pydict.items()}
async def _get_captured_rows_async(mock_write_client, expected_schema):
"""Helper to get all rows passed to append_rows."""
all_rows = []
for call in mock_write_client.append_rows.call_args_list:
requests_iter = call.args[0]
requests = []
if hasattr(requests_iter, "__aiter__"):
async for req in requests_iter:
requests.append(req)
else:
requests = list(requests_iter)
for request in requests:
# Parse the Arrow batch back to a dict for verification
try:
reader = pa.ipc.open_stream(
request.arrow_rows.rows.serialized_record_batch
)
table = reader.read_all()
except Exception:
# Fallback: try reading as a single batch
buf = pa.py_buffer(request.arrow_rows.rows.serialized_record_batch)
batch = pa.ipc.read_record_batch(buf, expected_schema)
table = pa.Table.from_batches([batch])
pydict = table.to_pylist()
all_rows.extend(pydict)
return all_rows
def _assert_common_fields(log_entry, event_type, agent="MyTestAgent"):
assert log_entry["event_type"] == event_type
assert log_entry["agent"] == agent
assert log_entry["session_id"] == "session-123"
assert log_entry["invocation_id"] == "inv-789"
def test_recursive_smart_truncate():
"""Test recursive smart truncate."""
obj = {
"a": "long string" * 10,
"b": ["short", "long string" * 10],
"c": {"d": "long string" * 10},
}
max_len = 10
truncated, is_truncated = (
bigquery_agent_analytics_plugin._recursive_smart_truncate(obj, max_len)
)
assert is_truncated
assert truncated["a"] == "long strin...[TRUNCATED]"
assert truncated["b"][0] == "short"
assert truncated["b"][1] == "long strin...[TRUNCATED]"
assert truncated["c"]["d"] == "long strin...[TRUNCATED]"
def test_recursive_smart_truncate_with_dataclasses():
"""Test recursive smart truncate with dataclasses."""
@dataclasses.dataclass
class LocalMissedKPI:
kpi: str
value: float
@dataclasses.dataclass
class LocalIncident:
id: str
kpi_missed: list[LocalMissedKPI]
status: str
incident = LocalIncident(
id="inc-123",
kpi_missed=[LocalMissedKPI(kpi="latency", value=99.9)],
status="active",
)
content = {"result": incident}
max_len = 1000
truncated, is_truncated = (
bigquery_agent_analytics_plugin._recursive_smart_truncate(
content, max_len
)
)
assert not is_truncated
assert isinstance(truncated["result"], dict)
assert truncated["result"]["id"] == "inc-123"
assert isinstance(truncated["result"]["kpi_missed"][0], dict)
assert truncated["result"]["kpi_missed"][0]["kpi"] == "latency"
class TestBigQueryAgentAnalyticsPlugin:
"""Tests for the BigQueryAgentAnalyticsPlugin."""
@pytest.mark.asyncio
async def test_plugin_disabled(
self,
mock_auth_default,
mock_bq_client,
mock_write_client,
invocation_context,
):
config = bigquery_agent_analytics_plugin.BigQueryLoggerConfig(enabled=False)
async with managed_plugin(
project_id=PROJECT_ID,
dataset_id=DATASET_ID,
table_id=TABLE_ID,
config=config,
) as plugin:
# user_message = types.Content(parts=[types.Part(text="Test")])
await plugin.on_user_message_callback(
invocation_context=invocation_context,
user_message=types.Content(parts=[types.Part(text="Test")]),
)
mock_auth_default.assert_not_called()
mock_bq_client.assert_not_called()
@pytest.mark.asyncio
async def test_enriched_metadata_logging(
self,
mock_auth_default,
mock_bq_client,
mock_write_client,
mock_to_arrow_schema,
dummy_arrow_schema,
callback_context,
):
# Setup
config = bigquery_agent_analytics_plugin.BigQueryLoggerConfig()
async with managed_plugin(PROJECT_ID, DATASET_ID, config=config) as plugin:
# Mock root agent
mock_root = mock.create_autospec(
base_agent.BaseAgent, instance=True, spec_set=True
)
type(mock_root).name = mock.PropertyMock(return_value="RootAgent")
callback_context._invocation_context.agent.root_agent = mock_root
# 1. Test root_agent_name and model extraction from request
llm_request = llm_request_lib.LlmRequest(
model="gemini-pro",
contents=[types.Content(parts=[types.Part(text="Hi")])],
)
await plugin.before_model_callback(
callback_context=callback_context, llm_request=llm_request
)
# 2. Test model_version and usage_metadata extraction from response
usage = types.GenerateContentResponseUsageMetadata(
prompt_token_count=10, candidates_token_count=20, total_token_count=30
)
llm_response = llm_response_lib.LlmResponse(
content=types.Content(parts=[types.Part(text="Hello")]),
usage_metadata=usage,
model_version="v1.2.3",
)
await plugin.after_model_callback(
callback_context=callback_context, llm_response=llm_response
)
# Verify captured rows from mock client
rows = await _get_captured_rows_async(mock_write_client, dummy_arrow_schema)
assert len(rows) == 2
# Check LLM_REQUEST row
# Sort by event_type to ensure consistent indexing
rows.sort(key=lambda x: x["event_type"])
request_row = rows[0] # LLM_REQUEST
response_row = rows[1] # LLM_RESPONSE
assert request_row["event_type"] == "LLM_REQUEST"
attr_req = json.loads(request_row["attributes"])
assert attr_req["root_agent_name"] == "RootAgent"
assert attr_req["model"] == "gemini-pro"
# Check LLM_RESPONSE row
assert response_row["event_type"] == "LLM_RESPONSE"
attr_res = json.loads(response_row["attributes"])
assert attr_res["root_agent_name"] == "RootAgent"
assert attr_res["model_version"] == "v1.2.3"
usage_meta = attr_res["usage_metadata"]
assert "prompt_token_count" in usage_meta
assert usage_meta["prompt_token_count"] == 10
mock_write_client.append_rows.assert_called()
@pytest.mark.asyncio
async def test_concurrent_span_management(
self,
mock_auth_default,
mock_bq_client,
mock_write_client,
mock_to_arrow_schema,
callback_context,
):
# Setup
config = bigquery_agent_analytics_plugin.BigQueryLoggerConfig()
plugin = bigquery_agent_analytics_plugin.BigQueryAgentAnalyticsPlugin(
PROJECT_ID, DATASET_ID, config=config
)
# Initialize trace in main context
bigquery_agent_analytics_plugin.TraceManager.init_trace(callback_context)
async def branch_1():
s_id = bigquery_agent_analytics_plugin.TraceManager.push_span(
callback_context, span_name="span-1"
)
await asyncio.sleep(0.02)
current_s_id = (
bigquery_agent_analytics_plugin.TraceManager.get_current_span_id()
)
assert s_id == current_s_id
bigquery_agent_analytics_plugin.TraceManager.pop_span()
return s_id
async def branch_2():
s_id = bigquery_agent_analytics_plugin.TraceManager.push_span(
callback_context, span_name="span-2"
)
await asyncio.sleep(0.02)
current_s_id = (
bigquery_agent_analytics_plugin.TraceManager.get_current_span_id()
)
assert s_id == current_s_id
bigquery_agent_analytics_plugin.TraceManager.pop_span()
return s_id
# Run concurrently
results = await asyncio.gather(branch_1(), branch_2())
# If they shared the same list/dict, they would interfere.
assert results[0] is not None
assert results[1] is not None
assert results[0] != results[1]
@pytest.mark.asyncio
async def test_event_allowlist(
self,
mock_write_client,
callback_context,
invocation_context,
mock_auth_default,
mock_bq_client,
mock_to_arrow_schema,
dummy_arrow_schema,
mock_asyncio_to_thread,
):
_ = mock_auth_default
_ = mock_bq_client
config = bigquery_agent_analytics_plugin.BigQueryLoggerConfig(
event_allowlist=["LLM_REQUEST"]
)
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()
llm_request = llm_request_lib.LlmRequest(
model="gemini-pro",
contents=[types.Content(parts=[types.Part(text="Prompt")])],
)
bigquery_agent_analytics_plugin.TraceManager.push_span(callback_context)
await plugin.before_model_callback(
callback_context=callback_context, llm_request=llm_request
)
await asyncio.sleep(0.01) # Allow background task to run
mock_write_client.append_rows.assert_called_once()
mock_write_client.append_rows.reset_mock()
user_message = types.Content(parts=[types.Part(text="What is up?")])
bigquery_agent_analytics_plugin.TraceManager.push_span(invocation_context)
await plugin.on_user_message_callback(
invocation_context=invocation_context, user_message=user_message
)
await asyncio.sleep(0.01) # Allow background task to run
mock_write_client.append_rows.assert_not_called()
@pytest.mark.asyncio
async def test_event_denylist(
self,
mock_write_client,
invocation_context,
mock_auth_default,
mock_bq_client,
mock_to_arrow_schema,
dummy_arrow_schema,
mock_asyncio_to_thread,
):
_ = mock_auth_default
_ = mock_bq_client
config = bigquery_agent_analytics_plugin.BigQueryLoggerConfig(
event_denylist=["USER_MESSAGE_RECEIVED"]
)
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()
user_message = types.Content(parts=[types.Part(text="What is up?")])
bigquery_agent_analytics_plugin.TraceManager.push_span(invocation_context)
await plugin.on_user_message_callback(
invocation_context=invocation_context, user_message=user_message
)
await asyncio.sleep(0.01)
mock_write_client.append_rows.assert_not_called()
bigquery_agent_analytics_plugin.TraceManager.push_span(invocation_context)
await plugin.before_run_callback(invocation_context=invocation_context)
await asyncio.sleep(0.01)
mock_write_client.append_rows.assert_called_once()
@pytest.mark.asyncio
async def test_content_formatter(
self,
mock_write_client,
invocation_context,
mock_auth_default,
mock_bq_client,
mock_to_arrow_schema,
dummy_arrow_schema,
mock_asyncio_to_thread,
):
"""Test content formatter."""
_ = mock_auth_default
_ = mock_bq_client
def redact_content(content, event_type):
return "[REDACTED]"
config = bigquery_agent_analytics_plugin.BigQueryLoggerConfig(
content_formatter=redact_content
)
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()
user_message = types.Content(parts=[types.Part(text="Secret message")])
bigquery_agent_analytics_plugin.TraceManager.push_span(invocation_context)
await plugin.on_user_message_callback(
invocation_context=invocation_context, user_message=user_message
)
await asyncio.sleep(0.01)
mock_write_client.append_rows.assert_called_once()
log_entry = await _get_captured_event_dict_async(
mock_write_client, dummy_arrow_schema
)
# If the formatter returns a string, it's stored directly.
assert log_entry["content"] == "[REDACTED]"
@pytest.mark.asyncio
async def test_content_formatter_error(
self,
mock_write_client,
invocation_context,
mock_auth_default,
mock_bq_client,
mock_to_arrow_schema,
dummy_arrow_schema,
mock_asyncio_to_thread,
):
"""Test content formatter error handling."""
_ = mock_auth_default
_ = mock_bq_client
def error_formatter(content, event_type):
raise ValueError("Formatter failed")
config = bigquery_agent_analytics_plugin.BigQueryLoggerConfig(
content_formatter=error_formatter
)
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()
user_message = types.Content(parts=[types.Part(text="Secret message")])
bigquery_agent_analytics_plugin.TraceManager.push_span(invocation_context)
await plugin.on_user_message_callback(
invocation_context=invocation_context, user_message=user_message
)
await asyncio.sleep(0.01)
mock_write_client.append_rows.assert_called_once()
log_entry = await _get_captured_event_dict_async(
mock_write_client, dummy_arrow_schema
)
# If formatter fails, it logs a warning and continues with original content.
assert log_entry["content"] == '{"text_summary": "Secret message"}'
@pytest.mark.asyncio
async def test_max_content_length(
self,
mock_write_client,
invocation_context,
callback_context,
mock_auth_default,
mock_bq_client,
mock_to_arrow_schema,
dummy_arrow_schema,
mock_asyncio_to_thread,
):
_ = mock_auth_default
_ = mock_bq_client
config = bigquery_agent_analytics_plugin.BigQueryLoggerConfig(
max_content_length=40
)
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()
# Test User Message Truncation
user_message = types.Content(
parts=[types.Part(text="12345678901234567890123456789012345678901")]
) # 41 chars
bigquery_agent_analytics_plugin.TraceManager.push_span(invocation_context)
await plugin.on_user_message_callback(
invocation_context=invocation_context, user_message=user_message
)
await asyncio.sleep(0.01)
mock_write_client.append_rows.assert_called_once()
log_entry = await _get_captured_event_dict_async(
mock_write_client, dummy_arrow_schema
)
assert (
log_entry["content"]
== '{"text_summary":'
' "1234567890123456789012345678901234567890...[TRUNCATED]"}'
)
assert log_entry["is_truncated"]
mock_write_client.append_rows.reset_mock()
# Test before_model_callback full content truncation
llm_request = llm_request_lib.LlmRequest(
model="gemini-pro",
config=types.GenerateContentConfig(
system_instruction=types.Content(
parts=[types.Part(text="System Instruction")]
)
),
contents=[
types.Content(role="user", parts=[types.Part(text="Prompt")])
],
)
bigquery_agent_analytics_plugin.TraceManager.push_span(callback_context)
await plugin.before_model_callback(
callback_context=callback_context, llm_request=llm_request
)
await asyncio.sleep(0.01)
mock_write_client.append_rows.assert_called_once()
log_entry = await _get_captured_event_dict_async(
mock_write_client, dummy_arrow_schema
)
# Full content: {"prompt": "text: 'Prompt'",
# "system_prompt": "text: 'System Instruction'"}
# In our new logic, we don't truncate the whole JSON string if it's valid JSON.
# Instead, we should have truncated the values within the dict, but currently we don't.
# For now, update test to reflect current behavior (valid JSON, no truncation of the whole string).
assert log_entry["content"].startswith(
'{"prompt": [{"role": "user", "content": "Prompt"}]'
)
assert log_entry["is_truncated"] is False
@pytest.mark.asyncio
async def test_max_content_length_tool_args(
self,
mock_write_client,
tool_context,
mock_auth_default,
mock_bq_client,
mock_to_arrow_schema,
dummy_arrow_schema,
mock_asyncio_to_thread,
):
_ = mock_auth_default
_ = mock_bq_client
_ = mock_to_arrow_schema
_ = mock_asyncio_to_thread
config = bigquery_agent_analytics_plugin.BigQueryLoggerConfig(
max_content_length=80
)
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()
mock_tool = mock.create_autospec(
base_tool_lib.BaseTool, instance=True, spec_set=True
)
type(mock_tool).name = mock.PropertyMock(return_value="MyTool")
type(mock_tool).description = mock.PropertyMock(
return_value="Description"
)
# Args length > 80
# {"param": "A" * 100} is > 100 chars.
bigquery_agent_analytics_plugin.TraceManager.push_span(tool_context)
await plugin.before_tool_callback(
tool=mock_tool,
tool_args={"param": "A" * 100},
tool_context=tool_context,
)
await asyncio.sleep(0.01)
mock_write_client.append_rows.assert_called_once()
log_entry = await _get_captured_event_dict_async(
mock_write_client, dummy_arrow_schema
)
_assert_common_fields(log_entry, "TOOL_STARTING")
# Now we do truncate nested values, and is_truncated flag is True
assert log_entry["is_truncated"]
content_dict = json.loads(log_entry["content"])
assert content_dict["tool"] == "MyTool"
assert content_dict["args"]["param"].endswith("...[TRUNCATED]")
@pytest.mark.asyncio
async def test_max_content_length_tool_args_no_truncation(
self,
mock_write_client,
tool_context,
mock_auth_default,
mock_bq_client,
mock_to_arrow_schema,
dummy_arrow_schema,
mock_asyncio_to_thread,
):
config = bigquery_agent_analytics_plugin.BigQueryLoggerConfig(
max_content_length=-1
)
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()
mock_tool = mock.create_autospec(
base_tool_lib.BaseTool, instance=True, spec_set=True
)
type(mock_tool).name = mock.PropertyMock(return_value="MyTool")
type(mock_tool).description = mock.PropertyMock(
return_value="Description"
)
# Args length > 80
# {"param": "A" * 100} is > 100 chars.
bigquery_agent_analytics_plugin.TraceManager.push_span(tool_context)
await plugin.before_tool_callback(
tool=mock_tool,
tool_args={"param": "A" * 100},
tool_context=tool_context,
)
await asyncio.sleep(0.01)
mock_write_client.append_rows.assert_called_once()
log_entry = await _get_captured_event_dict_async(
mock_write_client, dummy_arrow_schema
)
_assert_common_fields(log_entry, "TOOL_STARTING")
# No truncation
assert not log_entry["is_truncated"]
content_dict = json.loads(log_entry["content"])
assert content_dict["tool"] == "MyTool"
assert content_dict["args"]["param"] == "A" * 100
@pytest.mark.asyncio
async def test_max_content_length_tool_result(
self,
mock_write_client,
tool_context,
mock_auth_default,
mock_bq_client,
mock_asyncio_to_thread,
mock_to_arrow_schema,
dummy_arrow_schema,
):
"""Test max content length for tool result."""
_ = mock_auth_default
_ = mock_bq_client
_ = mock_to_arrow_schema
_ = mock_asyncio_to_thread
config = bigquery_agent_analytics_plugin.BigQueryLoggerConfig(
max_content_length=80
)
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()
mock_tool = mock.create_autospec(
base_tool_lib.BaseTool, instance=True, spec_set=True
)
type(mock_tool).name = mock.PropertyMock(return_value="MyTool")
# Result length > 80
# {"res": "A" * 100} is > 100 chars.
bigquery_agent_analytics_plugin.TraceManager.push_span(tool_context)
await plugin.after_tool_callback(
tool=mock_tool,
tool_args={},
tool_context=tool_context,
result={"res": "A" * 100},
)
await asyncio.sleep(0.01)
mock_write_client.append_rows.assert_called_once()
log_entry = await _get_captured_event_dict_async(
mock_write_client, dummy_arrow_schema
)
_assert_common_fields(log_entry, "TOOL_COMPLETED")
# Now we do truncate nested values, and is_truncated flag is True
assert log_entry["is_truncated"]
content_dict = json.loads(log_entry["content"])
assert content_dict["tool"] == "MyTool"
assert content_dict["result"]["res"].endswith("...[TRUNCATED]")
@pytest.mark.asyncio
async def test_max_content_length_tool_result_no_truncation(
self,
mock_write_client,
tool_context,
mock_auth_default,
mock_bq_client,
mock_to_arrow_schema,
dummy_arrow_schema,
mock_asyncio_to_thread,
):
"""Test max content length for tool result with no truncation."""
_ = mock_auth_default
_ = mock_bq_client
_ = mock_to_arrow_schema
_ = mock_asyncio_to_thread
config = bigquery_agent_analytics_plugin.BigQueryLoggerConfig(
max_content_length=-1
)
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()
mock_tool = mock.create_autospec(
base_tool_lib.BaseTool, instance=True, spec_set=True
)
type(mock_tool).name = mock.PropertyMock(return_value="MyTool")
# Result length > 80
# {"res": "A" * 100} is > 100 chars.
bigquery_agent_analytics_plugin.TraceManager.push_span(tool_context)
await plugin.after_tool_callback(
tool=mock_tool,
tool_args={},
tool_context=tool_context,
result={"res": "A" * 100},
)
await asyncio.sleep(0.01)
mock_write_client.append_rows.assert_called_once()
log_entry = await _get_captured_event_dict_async(
mock_write_client, dummy_arrow_schema
)
_assert_common_fields(log_entry, "TOOL_COMPLETED")
# No truncation
assert not log_entry["is_truncated"]
content_dict = json.loads(log_entry["content"])
assert content_dict["tool"] == "MyTool"
assert content_dict["result"]["res"] == "A" * 100
@pytest.mark.asyncio
async def test_max_content_length_tool_error(
self,
mock_write_client,
tool_context,
mock_auth_default,
mock_bq_client,
mock_to_arrow_schema,
dummy_arrow_schema,
mock_asyncio_to_thread,
):
config = bigquery_agent_analytics_plugin.BigQueryLoggerConfig(
max_content_length=80
)
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()
mock_tool = mock.create_autospec(
base_tool_lib.BaseTool, instance=True, spec_set=True
)
type(mock_tool).name = mock.PropertyMock(return_value="MyTool")
# Args length > 80
# {"arg": "A" * 100} is > 100 chars.
bigquery_agent_analytics_plugin.TraceManager.push_span(tool_context)
await plugin.on_tool_error_callback(
tool=mock_tool,
tool_args={"arg": "A" * 100},
tool_context=tool_context,
error=ValueError("Oops"),
)
await asyncio.sleep(0.01)
mock_write_client.append_rows.assert_called_once()
log_entry = await _get_captured_event_dict_async(
mock_write_client, dummy_arrow_schema
)
assert log_entry["content"].startswith(
'{"tool": "MyTool", "args": {"arg": "AAAAA'
)
# Check for truncation in the nested value
content_dict = json.loads(log_entry["content"])
assert content_dict["args"]["arg"].endswith("...[TRUNCATED]")
assert log_entry["is_truncated"]
assert log_entry["error_message"] == "Oops"
@pytest.mark.asyncio
async def test_on_user_message_callback_logs_correctly(
self,
bq_plugin_inst,
mock_write_client,
invocation_context,
dummy_arrow_schema,
):
user_message = types.Content(parts=[types.Part(text="What is up?")])
bigquery_agent_analytics_plugin.TraceManager.push_span(invocation_context)
await bq_plugin_inst.on_user_message_callback(
invocation_context=invocation_context, user_message=user_message
)
await asyncio.sleep(0.01)
log_entry = await _get_captured_event_dict_async(
mock_write_client, dummy_arrow_schema
)
_assert_common_fields(log_entry, "USER_MESSAGE_RECEIVED")
assert log_entry["content"] == '{"text_summary": "What is up?"}'
@pytest.mark.asyncio
async def test_offloading_with_connection_id(
self,
mock_write_client,
invocation_context,
mock_auth_default,
mock_bq_client,
mock_to_arrow_schema,
dummy_arrow_schema,
mock_asyncio_to_thread,
mock_storage_client,
):
_ = mock_auth_default
_ = mock_bq_client
_ = mock_to_arrow_schema
_ = mock_asyncio_to_thread
# Mock GCS bucket
mock_bucket = mock.Mock()
mock_blob = mock.Mock()
mock_bucket.blob.return_value = mock_blob
mock_bucket.name = "my-bucket"
mock_storage_client.return_value.bucket.return_value = mock_bucket
config = bigquery_agent_analytics_plugin.BigQueryLoggerConfig(
gcs_bucket_name="my-bucket",
connection_id="us.my-connection",
max_content_length=20, # Small limit to force offloading
)
async with managed_plugin(
PROJECT_ID, DATASET_ID, table_id=TABLE_ID, config=config
) as plugin:
await plugin._ensure_started(
storage_client=mock_storage_client.return_value
)
mock_write_client.append_rows.reset_mock()
# Create mixed content: one small inline, one large offloaded
small_text = "Small inline text"
large_text = "A" * 100
user_message = types.Content(
parts=[types.Part(text=small_text), types.Part(text=large_text)]
)
bigquery_agent_analytics_plugin.TraceManager.push_span(invocation_context)
await plugin.on_user_message_callback(
invocation_context=invocation_context, user_message=user_message
)
await asyncio.sleep(0.01)
mock_write_client.append_rows.assert_called_once()
log_entry = await _get_captured_event_dict_async(
mock_write_client, dummy_arrow_schema
)
# Verify content parts
assert len(log_entry["content_parts"]) == 2
# Part 0: Inline
part0 = log_entry["content_parts"][0]
assert part0["storage_mode"] == "INLINE"
assert part0["text"] == small_text
assert part0["object_ref"] is None
# Part 1: Offloaded
part1 = log_entry["content_parts"][1]
assert part1["storage_mode"] == "GCS_REFERENCE"
assert part1["uri"].startswith("gs://my-bucket/")
assert part1["object_ref"]["uri"] == part1["uri"]
assert part1["object_ref"]["authorizer"] == "us.my-connection"
assert json.loads(part1["object_ref"]["details"]) == {
"gcs_metadata": {"content_type": "text/plain"}
}
# Removed on_event_callback tests as they are no longer applicable in V2
@pytest.mark.asyncio
async def test_bigquery_client_initialization_failure(
self,
mock_auth_default,
mock_write_client,
invocation_context,
mock_asyncio_to_thread,
):
_ = mock_asyncio_to_thread
mock_auth_default.side_effect = auth_exceptions.GoogleAuthError(
"Auth failed"
)
async with managed_plugin(
project_id=PROJECT_ID,
dataset_id=DATASET_ID,
table_id=TABLE_ID,
) as plugin_with_fail:
with mock.patch(
"google.adk.plugins.bigquery_agent_analytics_plugin.logger"
) as mock_logger:
bigquery_agent_analytics_plugin.TraceManager.push_span(
invocation_context
)
await plugin_with_fail.on_user_message_callback(
invocation_context=invocation_context,
user_message=types.Content(parts=[types.Part(text="Test")]),
)
await asyncio.sleep(0.01)
mock_logger.error.assert_called_with(
"Failed to initialize BigQuery Plugin: %s", mock.ANY
)
mock_write_client.append_rows.assert_not_called()
@pytest.mark.asyncio
async def test_bigquery_insert_error_does_not_raise(
self, bq_plugin_inst, mock_write_client, invocation_context
):
_ = bq_plugin_inst
async def fake_append_rows_with_error(requests, **kwargs):
mock_append_rows_response = mock.MagicMock()
mock_append_rows_response.row_errors = [] # No row errors
mock_append_rows_response.error = mock.MagicMock()
mock_append_rows_response.error.code = 3 # INVALID_ARGUMENT
mock_append_rows_response.error.message = "Test BQ Error"
return _async_gen(mock_append_rows_response)
mock_write_client.append_rows.side_effect = fake_append_rows_with_error
with mock.patch(
"google.adk.plugins.bigquery_agent_analytics_plugin.logger"
) as mock_logger:
bigquery_agent_analytics_plugin.TraceManager.push_span(invocation_context)
await bq_plugin_inst.on_user_message_callback(
invocation_context=invocation_context,
user_message=types.Content(parts=[types.Part(text="Test")]),
)
await asyncio.sleep(0.01)
# The logger is called multiple times, check that one of them is the error message
# Or just check that it was called with the expected message at some point
mock_logger.error.assert_any_call(
"Non-retryable BigQuery error: %s", "Test BQ Error"
)
mock_write_client.append_rows.assert_called_once()
@pytest.mark.asyncio
async def test_bigquery_insert_retryable_error(
self, bq_plugin_inst, mock_write_client, invocation_context
):
"""Test that retryable BigQuery errors are logged and retried."""
async def fake_append_rows_with_retryable_error(requests, **kwargs):
mock_append_rows_response = mock.MagicMock()
mock_append_rows_response.row_errors = [] # No row errors
mock_append_rows_response.error = mock.MagicMock()
mock_append_rows_response.error.code = 10 # ABORTED (retryable)
mock_append_rows_response.error.message = "Test BQ Retryable Error"
return _async_gen(mock_append_rows_response)
mock_write_client.append_rows.side_effect = (
fake_append_rows_with_retryable_error
)
with mock.patch(
"google.adk.plugins.bigquery_agent_analytics_plugin.logger"
) as mock_logger:
bigquery_agent_analytics_plugin.TraceManager.push_span(invocation_context)
await bq_plugin_inst.on_user_message_callback(
invocation_context=invocation_context,
user_message=types.Content(parts=[types.Part(text="Test")]),
)
await asyncio.sleep(0.01)
mock_logger.warning.assert_any_call(
"BigQuery Write API returned error code %s: %s",
10,
"Test BQ Retryable Error",
)
# Should be called at least once. Retries are hard to test due to async backoff.
assert mock_write_client.append_rows.call_count >= 1
@pytest.mark.asyncio
async def test_schema_mismatch_error_handling(
self, bq_plugin_inst, mock_write_client, invocation_context
):
async def fake_append_rows_with_schema_error(requests, **kwargs):
mock_resp = mock.MagicMock()
mock_resp.row_errors = []
mock_resp.error = mock.MagicMock()
mock_resp.error.code = 3
mock_resp.error.message = (
"Schema mismatch: Field 'new_field' not found in table."
)
return _async_gen(mock_resp)
mock_write_client.append_rows.side_effect = (
fake_append_rows_with_schema_error
)
with mock.patch(
"google.adk.plugins.bigquery_agent_analytics_plugin.logger"
) as mock_logger:
bigquery_agent_analytics_plugin.TraceManager.push_span(invocation_context)
await bq_plugin_inst.on_user_message_callback(
invocation_context=invocation_context,
user_message=types.Content(parts=[types.Part(text="Test")]),
)
await asyncio.sleep(0.01)
mock_logger.error.assert_called_with(
"BigQuery Schema Mismatch: %s. This usually means the"
" table schema does not match the expected schema.",
"Schema mismatch: Field 'new_field' not found in table.",
)
@pytest.mark.asyncio
async def test_close(self, bq_plugin_inst, mock_bq_client, mock_write_client):
"""Test plugin shutdown."""
await bq_plugin_inst.shutdown()
# shutdown calls transport.close() on all clients
assert mock_write_client.transport.close.call_count >= 1
# Verify loop states are cleared
assert not bq_plugin_inst._loop_state_by_loop
@pytest.mark.asyncio
async def test_before_run_callback_logs_correctly(
self,
bq_plugin_inst,
mock_write_client,
invocation_context,
dummy_arrow_schema,
):
"""Test before_run_callback logs correctly."""
bigquery_agent_analytics_plugin.TraceManager.push_span(invocation_context)
await bq_plugin_inst.before_run_callback(
invocation_context=invocation_context
)
await asyncio.sleep(0.01)
log_entry = await _get_captured_event_dict_async(
mock_write_client, dummy_arrow_schema
)
_assert_common_fields(log_entry, "INVOCATION_STARTING")
assert log_entry["content"] is None
@pytest.mark.asyncio
async def test_after_run_callback_logs_correctly(
self,
bq_plugin_inst,
mock_write_client,
invocation_context,
dummy_arrow_schema,
):
bigquery_agent_analytics_plugin.TraceManager.push_span(invocation_context)
await bq_plugin_inst.after_run_callback(
invocation_context=invocation_context
)
await asyncio.sleep(0.01)
log_entry = await _get_captured_event_dict_async(
mock_write_client, dummy_arrow_schema
)
_assert_common_fields(log_entry, "INVOCATION_COMPLETED")
assert log_entry["content"] is None
@pytest.mark.asyncio
async def test_before_agent_callback_logs_correctly(
self,
bq_plugin_inst,
mock_write_client,
mock_agent,
callback_context,
dummy_arrow_schema,
):
bigquery_agent_analytics_plugin.TraceManager.push_span(callback_context)
await bq_plugin_inst.before_agent_callback(
agent=mock_agent, callback_context=callback_context
)
await asyncio.sleep(0.01)
log_entry = await _get_captured_event_dict_async(
mock_write_client, dummy_arrow_schema
)
_assert_common_fields(log_entry, "AGENT_STARTING")
assert log_entry["content"] == "Test Instruction"
@pytest.mark.asyncio
async def test_after_agent_callback_logs_correctly(
self,
bq_plugin_inst,
mock_write_client,
mock_agent,
callback_context,
dummy_arrow_schema,
):
bigquery_agent_analytics_plugin.TraceManager.push_span(callback_context)
await bq_plugin_inst.after_agent_callback(
agent=mock_agent, callback_context=callback_context
)
await asyncio.sleep(0.01)
log_entry = await _get_captured_event_dict_async(
mock_write_client, dummy_arrow_schema
)
_assert_common_fields(log_entry, "AGENT_COMPLETED")
assert log_entry["content"] is None
# Latency should be an int >= 0 now that we instrument it
assert log_entry["latency_ms"] is not None
latency_dict = json.loads(log_entry["latency_ms"])
assert latency_dict["total_ms"] >= 0
@pytest.mark.asyncio
async def test_before_model_callback_logs_correctly(
self,
bq_plugin_inst,
mock_write_client,
callback_context,
dummy_arrow_schema,
):
llm_request = llm_request_lib.LlmRequest(
model="gemini-pro",
contents=[
types.Content(role="user", parts=[types.Part(text="Prompt")])
],
)
bigquery_agent_analytics_plugin.TraceManager.push_span(callback_context)
await bq_plugin_inst.before_model_callback(
callback_context=callback_context, llm_request=llm_request
)
await asyncio.sleep(0.01)
log_entry = await _get_captured_event_dict_async(
mock_write_client, dummy_arrow_schema
)
_assert_common_fields(log_entry, "LLM_REQUEST")
assert "Prompt" in log_entry["content"]
@pytest.mark.asyncio
async def test_before_model_callback_with_params_and_tools(
self,
bq_plugin_inst,
mock_write_client,
callback_context,
dummy_arrow_schema,
):
llm_request = llm_request_lib.LlmRequest(
model="gemini-pro",
config=types.GenerateContentConfig(
temperature=0.5,
top_p=0.9,
system_instruction=types.Content(parts=[types.Part(text="Sys")]),
),
contents=[types.Content(role="user", parts=[types.Part(text="User")])],
)
# Manually set tools_dict as it is excluded from init
llm_request.tools_dict = {"tool1": "func1", "tool2": "func2"}
bigquery_agent_analytics_plugin.TraceManager.push_span(callback_context)
await bq_plugin_inst.before_model_callback(
callback_context=callback_context, llm_request=llm_request
)
await asyncio.sleep(0.01)
log_entry = await _get_captured_event_dict_async(
mock_write_client, dummy_arrow_schema
)
_assert_common_fields(log_entry, "LLM_REQUEST")
# Verify content is JSON and has correct fields
assert "content" in log_entry
content_dict = json.loads(log_entry["content"])
assert content_dict["prompt"] == [{"role": "user", "content": "User"}]
assert content_dict["system_prompt"] == "Sys"
# Verify attributes
assert "attributes" in log_entry
attributes = json.loads(log_entry["attributes"])
assert attributes["llm_config"]["temperature"] == 0.5
assert attributes["llm_config"]["top_p"] == 0.9
assert attributes["llm_config"]["top_p"] == 0.9
assert attributes["tools"] == ["tool1", "tool2"]
@pytest.mark.asyncio
async def test_before_model_callback_with_full_config(
self,
bq_plugin_inst,
mock_write_client,
callback_context,
dummy_arrow_schema,
):
"""Test that all config fields, including falsy values and labels, are logged."""
llm_request = llm_request_lib.LlmRequest(
model="gemini-pro",
config=types.GenerateContentConfig(
temperature=0.0,
top_p=0.1,
top_k=5.0,
candidate_count=5,
max_output_tokens=65000,
stop_sequences=["STOP"],
presence_penalty=0.1,
frequency_penalty=0.5,
seed=42,
response_logprobs=True,
logprobs=3,
labels={"llm.agent.name": "test_agent"},
),
contents=[types.Content(role="user", parts=[types.Part(text="User")])],
)
bigquery_agent_analytics_plugin.TraceManager.push_span(callback_context)
await bq_plugin_inst.before_model_callback(
callback_context=callback_context, llm_request=llm_request
)
await asyncio.sleep(0.01)
log_entry = await _get_captured_event_dict_async(
mock_write_client, dummy_arrow_schema
)
_assert_common_fields(log_entry, "LLM_REQUEST")
# Verify attributes
assert "attributes" in log_entry
attributes = json.loads(log_entry["attributes"])
llm_config = attributes.get("llm_config", {})
expected_llm_config = {
"temperature": 0.0,
"top_p": 0.1,
"top_k": 5.0,
"candidate_count": 5,
"max_output_tokens": 65000,
"stop_sequences": ["STOP"],
"presence_penalty": 0.1,
"frequency_penalty": 0.5,
"seed": 42,
"response_logprobs": True,
"logprobs": 3,
}
assert llm_config == expected_llm_config
assert attributes.get("labels") == {"llm.agent.name": "test_agent"}
@pytest.mark.asyncio
async def test_before_model_callback_multipart_separator(
self,
bq_plugin_inst,
mock_write_client,
callback_context,
dummy_arrow_schema,
):
llm_request = llm_request_lib.LlmRequest(
model="gemini-pro",
contents=[
types.Content(
role="user",
parts=[types.Part(text="Part1"), types.Part(text="Part2")],
)
],
)
bigquery_agent_analytics_plugin.TraceManager.push_span(callback_context)
await bq_plugin_inst.before_model_callback(
callback_context=callback_context, llm_request=llm_request
)
await asyncio.sleep(0.01)
log_entry = await _get_captured_event_dict_async(
mock_write_client, dummy_arrow_schema
)
content_dict = json.loads(log_entry["content"])
# Verify the separator is " | "
assert content_dict["prompt"][0]["content"] == "Part1 | Part2"
@pytest.mark.asyncio
async def test_after_model_callback_text_response(
self,
bq_plugin_inst,
mock_write_client,
callback_context,
dummy_arrow_schema,
):
llm_response = llm_response_lib.LlmResponse(
content=types.Content(parts=[types.Part(text="Model response")]),
usage_metadata=types.UsageMetadata(
prompt_token_count=10, total_token_count=15
),
)
bigquery_agent_analytics_plugin.TraceManager.push_span(callback_context)
await bq_plugin_inst.after_model_callback(
callback_context=callback_context,
llm_response=llm_response,
# latency_ms is now calculated internally via TraceManager
)
await asyncio.sleep(0.01)
log_entry = await _get_captured_event_dict_async(
mock_write_client, dummy_arrow_schema
)
_assert_common_fields(log_entry, "LLM_RESPONSE")
content_dict = json.loads(log_entry["content"])
assert content_dict["response"] == "text: 'Model response'"
assert content_dict["usage"]["prompt"] == 10
assert content_dict["usage"]["total"] == 15
assert log_entry["error_message"] is None
latency_dict = json.loads(log_entry["latency_ms"])
# Latency comes from time.time(), so we can't assert exact 100ms
# But it should be present
assert latency_dict["total_ms"] >= 0
# tfft is passed via kwargs if present, or we can mock it.
# In this test we didn't pass it in kwargs in the updated call above, so it might be missing unless we add it back to kwargs.
# The original test passed it as kwarg.
@pytest.mark.asyncio
async def test_after_model_callback_tool_call(
self,
bq_plugin_inst,
mock_write_client,
callback_context,
dummy_arrow_schema,
):
tool_fc = types.FunctionCall(name="get_weather", args={"location": "Paris"})
llm_response = llm_response_lib.LlmResponse(
content=types.Content(parts=[types.Part(function_call=tool_fc)]),
usage_metadata=types.UsageMetadata(
prompt_token_count=10, total_token_count=15
),
)
bigquery_agent_analytics_plugin.TraceManager.push_span(callback_context)
await bq_plugin_inst.after_model_callback(
callback_context=callback_context,
llm_response=llm_response,
)
await asyncio.sleep(0.01)
log_entry = await _get_captured_event_dict_async(
mock_write_client, dummy_arrow_schema
)
_assert_common_fields(log_entry, "LLM_RESPONSE")
content_dict = json.loads(log_entry["content"])
assert content_dict["response"] == "call: get_weather"
assert content_dict["usage"]["prompt"] == 10
assert content_dict["usage"]["total"] == 15
assert log_entry["error_message"] is None
@pytest.mark.asyncio
async def test_before_tool_callback_logs_correctly(
self, bq_plugin_inst, mock_write_client, tool_context, dummy_arrow_schema
):
mock_tool = mock.create_autospec(
base_tool_lib.BaseTool, instance=True, spec_set=True
)
type(mock_tool).name = mock.PropertyMock(return_value="MyTool")
type(mock_tool).description = mock.PropertyMock(return_value="Description")
bigquery_agent_analytics_plugin.TraceManager.push_span(tool_context)
await bq_plugin_inst.before_tool_callback(
tool=mock_tool, tool_args={"param": "value"}, tool_context=tool_context
)
await asyncio.sleep(0.01)
log_entry = await _get_captured_event_dict_async(
mock_write_client, dummy_arrow_schema
)
_assert_common_fields(log_entry, "TOOL_STARTING")
content_dict = json.loads(log_entry["content"])
assert content_dict["tool"] == "MyTool"
assert content_dict["args"] == {"param": "value"}
@pytest.mark.asyncio
async def test_after_tool_callback_logs_correctly(
self, bq_plugin_inst, mock_write_client, tool_context, dummy_arrow_schema
):
mock_tool = mock.create_autospec(
base_tool_lib.BaseTool, instance=True, spec_set=True
)
type(mock_tool).name = mock.PropertyMock(return_value="MyTool")
type(mock_tool).description = mock.PropertyMock(return_value="Description")
bigquery_agent_analytics_plugin.TraceManager.push_span(tool_context)
await bq_plugin_inst.after_tool_callback(
tool=mock_tool,
tool_args={"arg1": "val1"},
tool_context=tool_context,
result={"res": "success"},
)
await asyncio.sleep(0.01)
log_entry = await _get_captured_event_dict_async(
mock_write_client, dummy_arrow_schema
)
_assert_common_fields(log_entry, "TOOL_COMPLETED")
content_dict = json.loads(log_entry["content"])
assert content_dict["tool"] == "MyTool"
assert content_dict["result"] == {"res": "success"}
@pytest.mark.asyncio
async def test_after_tool_callback_no_state_delta_logging(
self, bq_plugin_inst, mock_write_client, tool_context, dummy_arrow_schema
):
"""State deltas are now logged via on_event_callback, not after_tool."""
mock_tool = mock.create_autospec(
base_tool_lib.BaseTool, instance=True, spec_set=True
)
type(mock_tool).name = mock.PropertyMock(return_value="StateTool")
type(mock_tool).description = mock.PropertyMock(return_value="Sets state")
# Simulate a tool modifying the state
tool_context.actions.state_delta["new_key"] = "new_value"
bigquery_agent_analytics_plugin.TraceManager.push_span(tool_context)
await bq_plugin_inst.after_tool_callback(
tool=mock_tool,
tool_args={"arg1": "val1"},
tool_context=tool_context,
result={"res": "success"},
)
await asyncio.sleep(0.01)
# Only TOOL_COMPLETED should be logged; STATE_DELTA is handled
# by on_event_callback now.
rows = await _get_captured_rows_async(mock_write_client, dummy_arrow_schema)
assert len(rows) == 1
assert rows[0]["event_type"] == "TOOL_COMPLETED"
@pytest.mark.asyncio
async def test_on_event_callback_logs_state_delta(
self,
bq_plugin_inst,
mock_write_client,
invocation_context,
dummy_arrow_schema,
):
"""on_event_callback logs STATE_DELTA for events with state changes."""
state_delta = {"key": "value", "new_key": 123}
event = event_lib.Event(
author="test_agent",
actions=event_actions_lib.EventActions(state_delta=state_delta),
)
bigquery_agent_analytics_plugin.TraceManager.push_span(invocation_context)
result = await bq_plugin_inst.on_event_callback(
invocation_context=invocation_context, event=event
)
# Must return None to not modify the event
assert result is None
await asyncio.sleep(0.01)
log_entry = await _get_captured_event_dict_async(
mock_write_client, dummy_arrow_schema
)
_assert_common_fields(log_entry, "STATE_DELTA")
assert log_entry["content"] is None
attributes = json.loads(log_entry["attributes"])
assert attributes["state_delta"] == state_delta
@pytest.mark.asyncio
async def test_on_event_callback_ignores_empty_state_delta(
self,
bq_plugin_inst,
mock_write_client,
invocation_context,
):
"""on_event_callback should not log when state_delta is empty."""
event = event_lib.Event(
author="test_agent",
actions=event_actions_lib.EventActions(state_delta={}),
)
result = await bq_plugin_inst.on_event_callback(
invocation_context=invocation_context, event=event
)
assert result is None
# No events should have been logged
mock_write_client.append_rows.assert_not_called()
@pytest.mark.asyncio
async def test_log_event_with_session_metadata(
self,
bq_plugin_inst,
mock_write_client,
callback_context,
dummy_arrow_schema,
):
"""Test that session metadata is logged when enabled."""
# Setup session state with user metadata
session = callback_context._invocation_context.session
type(session).state = mock.PropertyMock(
return_value={"thread_id": "gchat-123", "customer_id": "cust-42"}
)
# Ensure config enabled (default is True)
bq_plugin_inst.config.log_session_metadata = True
await bq_plugin_inst._log_event(
"TEST_EVENT",
callback_context,
raw_content="test content",
)
await asyncio.sleep(0.01)
log_entry = await _get_captured_event_dict_async(
mock_write_client, dummy_arrow_schema
)
attributes = json.loads(log_entry["attributes"])
meta = attributes["session_metadata"]
assert meta["session_id"] == session.id
assert meta["app_name"] == session.app_name
assert meta["user_id"] == session.user_id
assert meta["state"] == {
"thread_id": "gchat-123",
"customer_id": "cust-42",
}
@pytest.mark.asyncio
async def test_log_event_with_custom_tags(
self,
bq_plugin_inst,
mock_write_client,
callback_context,
dummy_arrow_schema,
):
"""Test that custom tags are logged."""
custom_tags = {"agent_role": "sales", "env": "prod"}
bq_plugin_inst.config.custom_tags = custom_tags
await bq_plugin_inst._log_event(
"TEST_EVENT",
callback_context,
raw_content="test content",
)
await asyncio.sleep(0.01)
log_entry = await _get_captured_event_dict_async(
mock_write_client, dummy_arrow_schema
)
attributes = json.loads(log_entry["attributes"])
assert attributes["custom_tags"] == custom_tags
@pytest.mark.asyncio
async def test_on_model_error_callback_logs_correctly(
self,
bq_plugin_inst,
mock_write_client,
callback_context,
dummy_arrow_schema,
):
llm_request = llm_request_lib.LlmRequest(
model="gemini-pro",
contents=[types.Content(parts=[types.Part(text="Prompt")])],
)
error = ValueError("LLM failed")
bigquery_agent_analytics_plugin.TraceManager.push_span(callback_context)
await bq_plugin_inst.on_model_error_callback(
callback_context=callback_context, llm_request=llm_request, error=error
)
await asyncio.sleep(0.01)
log_entry = await _get_captured_event_dict_async(
mock_write_client, dummy_arrow_schema
)
_assert_common_fields(log_entry, "LLM_ERROR")
assert log_entry["content"] is None
assert log_entry["error_message"] == "LLM failed"
@pytest.mark.asyncio
async def test_on_tool_error_callback_logs_correctly(
self, bq_plugin_inst, mock_write_client, tool_context, dummy_arrow_schema
):
mock_tool = mock.create_autospec(
base_tool_lib.BaseTool, instance=True, spec_set=True
)
type(mock_tool).name = mock.PropertyMock(return_value="MyTool")
type(mock_tool).description = mock.PropertyMock(return_value="Description")
error = TimeoutError("Tool timed out")
bigquery_agent_analytics_plugin.TraceManager.push_span(tool_context)
await bq_plugin_inst.on_tool_error_callback(
tool=mock_tool,
tool_args={"param": "value"},
tool_context=tool_context,
error=error,
)
await asyncio.sleep(0.01)
log_entry = await _get_captured_event_dict_async(
mock_write_client, dummy_arrow_schema
)
_assert_common_fields(log_entry, "TOOL_ERROR")
content_dict = json.loads(log_entry["content"])
assert content_dict["tool"] == "MyTool"
assert content_dict["args"] == {"param": "value"}
assert log_entry["error_message"] == "Tool timed out"
@pytest.mark.asyncio
async def test_table_creation_options(
self,
mock_auth_default,
mock_bq_client,
mock_write_client,
mock_to_arrow_schema,
mock_asyncio_to_thread,
):
async with managed_plugin(
PROJECT_ID, DATASET_ID, table_id=TABLE_ID
) as plugin:
mock_bq_client.get_table.side_effect = cloud_exceptions.NotFound(
"Not found"
)
await plugin._ensure_started()
# Verify create_table was called with correct table options
mock_bq_client.create_table.assert_called_once()
call_args = mock_bq_client.create_table.call_args
table_arg = call_args[0][0]
assert isinstance(table_arg, bigquery.Table)
assert table_arg.time_partitioning.type_ == "DAY"
assert table_arg.time_partitioning.field == "timestamp"
assert table_arg.clustering_fields == ["event_type", "agent", "user_id"]
# Verify schema descriptions are present (spot check)
timestamp_field = next(
f for f in table_arg.schema if f.name == "timestamp"
)
assert (
timestamp_field.description
== "The UTC timestamp when the event occurred. Used for ordering"
" events"
" within a session."
)
@pytest.mark.asyncio
async def test_init_in_thread_pool(
self,
mock_auth_default,
mock_bq_client,
mock_write_client,
mock_to_arrow_schema,
mock_asyncio_to_thread,
invocation_context,
):
"""Verifies that the plugin can be initialized from a thread pool."""
async with managed_plugin(
PROJECT_ID, DATASET_ID, table_id=TABLE_ID
) as plugin:
def _run_in_thread(p):
# In a real thread pool, there might not be an event loop.
# However, since we are calling an async method (_ensure_started),
# we must run it in an event loop. The issue was that _lazy_setup
# called get_event_loop() which fails in threads without a loop.
# Here we simulate the condition by running in a thread and creating a new loop if needed,
# but the key is that the plugin's internal calls should use the correct loop.
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
try:
# _ensure_started is called by managed_plugin, but we need to ensure
# that if it were called in a thread, it would work.
# For this test, we just ensure the plugin is accessible and started.
loop.run_until_complete(p._ensure_started())
finally:
loop.close()
# Run in a separate thread to simulate ThreadPoolExecutor-0_0
from concurrent.futures import ThreadPoolExecutor
with ThreadPoolExecutor(max_workers=1) as executor:
future = executor.submit(_run_in_thread, plugin)
future.result() # Should not raise "no current event loop"
assert plugin._started
# Verify loop states are populated
assert plugin._loop_state_by_loop
@pytest.mark.asyncio
async def test_multimodal_offloading(
self,
mock_write_client,
callback_context,
mock_auth_default,
mock_bq_client,
mock_to_arrow_schema,
dummy_arrow_schema,
mock_storage_client,
):
# Setup
bucket_name = "test-bucket"
config = bigquery_agent_analytics_plugin.BigQueryLoggerConfig(
gcs_bucket_name=bucket_name
)
async with managed_plugin(
PROJECT_ID, DATASET_ID, table_id=TABLE_ID, config=config
) as plugin:
await plugin._ensure_started(
storage_client=mock_storage_client.return_value
)
# Mock GCS bucket and blob
mock_bucket = mock_storage_client.return_value.bucket.return_value
mock_bucket.name = bucket_name
mock_blob = mock_bucket.blob.return_value
# Create content with large text that should be offloaded
large_text = "A" * (32 * 1024 + 1)
llm_request = llm_request_lib.LlmRequest(
model="gemini-pro",
contents=[types.Content(parts=[types.Part(text=large_text)])],
)
# Execute
await plugin.before_model_callback(
callback_context=callback_context, llm_request=llm_request
)
# Use flush instead of sleep for robustness
await plugin.flush()
# Verify GCS upload
mock_blob.upload_from_string.assert_called_once()
args, kwargs = mock_blob.upload_from_string.call_args
assert args[0] == large_text
assert kwargs["content_type"] == "text/plain"
# Verify BQ write
mock_write_client.append_rows.assert_called_once()
event_dict = await _get_captured_event_dict_async(
mock_write_client, dummy_arrow_schema
)
content_parts = event_dict["content_parts"]
assert len(content_parts) == 1
assert content_parts[0]["storage_mode"] == "GCS_REFERENCE"
assert content_parts[0]["uri"].startswith(f"gs://{bucket_name}/")
@pytest.mark.asyncio
async def test_quota_project_id_used_in_client(
self,
mock_bq_client,
mock_to_arrow_schema,
mock_asyncio_to_thread,
):
mock_creds = mock.create_autospec(
google.auth.credentials.Credentials, instance=True, spec_set=True
)
mock_creds.quota_project_id = "quota-project"
with mock.patch.object(
google.auth,
"default",
autospec=True,
return_value=(mock_creds, PROJECT_ID),
) as mock_auth_default:
with mock.patch.object(
bigquery_agent_analytics_plugin,
"BigQueryWriteAsyncClient",
autospec=True,
) as mock_bq_write_cls:
async with managed_plugin(
project_id=PROJECT_ID,
dataset_id=DATASET_ID,
table_id=TABLE_ID,
) as plugin:
await plugin._ensure_started()
mock_auth_default.assert_called_once()
mock_bq_write_cls.assert_called_once()
_, kwargs = mock_bq_write_cls.call_args
assert kwargs["client_options"].quota_project_id == "quota-project"
@pytest.mark.asyncio
async def test_no_quota_project_when_creds_lack_it(
self,
mock_bq_client,
mock_to_arrow_schema,
mock_asyncio_to_thread,
):
"""Verify no quota_project_id is set when credentials don't provide one.
This is critical for Workload Identity Federation flows where setting
quota_project_id on the client breaks auth token refresh (issue #4370).
"""
mock_creds = mock.create_autospec(
google.auth.credentials.Credentials, instance=True, spec_set=True
)
mock_creds.quota_project_id = None
with mock.patch.object(
google.auth,
"default",
autospec=True,
return_value=(mock_creds, PROJECT_ID),
):
with mock.patch.object(
bigquery_agent_analytics_plugin,
"BigQueryWriteAsyncClient",
autospec=True,
) as mock_bq_write_cls:
async with managed_plugin(
project_id=PROJECT_ID,
dataset_id=DATASET_ID,
table_id=TABLE_ID,
) as plugin:
await plugin._ensure_started()
mock_bq_write_cls.assert_called_once()
_, kwargs = mock_bq_write_cls.call_args
assert kwargs["client_options"] is None
@pytest.mark.asyncio
async def test_pickle_safety(self, mock_auth_default, mock_bq_client):
"""Test that the plugin can be pickled safely."""
import pickle
config = bigquery_agent_analytics_plugin.BigQueryLoggerConfig(enabled=True)
plugin = bigquery_agent_analytics_plugin.BigQueryAgentAnalyticsPlugin(
PROJECT_ID, DATASET_ID, table_id=TABLE_ID, config=config
)
# Test pickling before start
pickled = pickle.dumps(plugin)
unpickled = pickle.loads(pickled)
assert unpickled.project_id == PROJECT_ID
assert unpickled._setup_lock is None
assert unpickled._executor is None
# Start the plugin
await plugin._ensure_started()
assert plugin._executor is not None
try:
# Test pickling after start
pickled_started = pickle.dumps(plugin)
unpickled_started = pickle.loads(pickled_started)
assert unpickled_started.project_id == PROJECT_ID
# Runtime objects should be None after unpickling
assert unpickled_started._setup_lock is None
assert unpickled_started._executor is None
assert not unpickled_started._loop_state_by_loop
finally:
await plugin.shutdown()
@pytest.mark.asyncio
async def test_span_hierarchy_llm_call(
self,
bq_plugin_inst,
mock_write_client,
callback_context,
dummy_arrow_schema,
):
"""Verifies that LLM events have correct Span ID hierarchy."""
# 1. Start Agent Span
bigquery_agent_analytics_plugin.TraceManager.push_span(callback_context)
_, _ = (
bigquery_agent_analytics_plugin.TraceManager.get_current_span_and_parent()
)
agent_span_id = (
bigquery_agent_analytics_plugin.TraceManager.get_current_span_id()
)
# 2. Start LLM Span (Implicitly handled if we push it?
# Actually before_model_callback assumes a span is pushed for the LLM call if we want one?
# No, usually the Runner/Agent pushes a span BEFORE calling before_model_callback?
# Let's verify usage in agent.py or plugin.
# Plugin does NOT push spans automatically for LLM. It relies on TraceManager being managed externally
# OR it uses current span.
# Wait, the Runner pushes spans.
# 3. LLM Request
llm_request = llm_request_lib.LlmRequest(
model="gemini-pro",
contents=[types.Content(parts=[types.Part(text="Prompt")])],
)
await bq_plugin_inst.before_model_callback(
callback_context=callback_context, llm_request=llm_request
)
await asyncio.sleep(0.01)
# Capture the actual LLM Span ID (pushed by before_model_callback)
llm_span_id = (
bigquery_agent_analytics_plugin.TraceManager.get_current_span_id()
)
# Now that we push a new span for LLM calls, it should differ from agent_span_id
assert llm_span_id != agent_span_id
log_entry_req = await _get_captured_event_dict_async(
mock_write_client, dummy_arrow_schema
)
assert log_entry_req["event_type"] == "LLM_REQUEST"
assert log_entry_req["span_id"] == llm_span_id
# The parent of the LLM span should be the Agent span
assert log_entry_req["parent_span_id"] == agent_span_id
mock_write_client.append_rows.reset_mock()
# 4. LLM Response
# In the actual flow, after_model_callback pops the span.
# But explicitly via TraceManager.pop_span()?
# No, after_model_callback calls TraceManager.pop_span().
# So we should validly call it.
llm_response = llm_response_lib.LlmResponse(
content=types.Content(parts=[types.Part(text="Response")]),
)
await bq_plugin_inst.after_model_callback(
callback_context=callback_context, llm_response=llm_response
)
await asyncio.sleep(0.01)
log_entry_resp = await _get_captured_event_dict_async(
mock_write_client, dummy_arrow_schema
)
assert log_entry_resp["event_type"] == "LLM_RESPONSE"
assert log_entry_resp["span_id"] == llm_span_id
# The parent of the LLM span should be the Agent span
assert log_entry_resp["parent_span_id"] == agent_span_id
# Verify LLM Span was popped and we are back to Agent Span
assert (
bigquery_agent_analytics_plugin.TraceManager.get_current_span_id()
== agent_span_id
)
# Clean up Agent Span
bigquery_agent_analytics_plugin.TraceManager.pop_span()
assert (
not bigquery_agent_analytics_plugin.TraceManager.get_current_span_id()
)
@pytest.mark.asyncio
async def test_custom_object_serialization(
self,
mock_write_client,
tool_context,
mock_auth_default,
mock_bq_client,
mock_to_arrow_schema,
dummy_arrow_schema,
mock_asyncio_to_thread,
):
"""Verifies that custom objects (Dataclasses) are serialized to dicts."""
_ = mock_auth_default
_ = mock_bq_client
@dataclasses.dataclass
class LocalMissedKPI:
kpi: str
value: float
@dataclasses.dataclass
class LocalIncident:
id: str
kpi_missed: list[LocalMissedKPI]
status: str
incident = LocalIncident(
id="inc-123",
kpi_missed=[LocalMissedKPI(kpi="latency", value=99.9)],
status="active",
)
config = bigquery_agent_analytics_plugin.BigQueryLoggerConfig()
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()
content = {"result": incident}
# Verify full flow
await plugin._log_event(
"TOOL_PARTIAL",
tool_context,
raw_content=content,
)
await asyncio.sleep(0.01)
mock_write_client.append_rows.assert_called_once()
log_entry = await _get_captured_event_dict_async(
mock_write_client, dummy_arrow_schema
)
# Content should be valid JSON string
content_json = json.loads(log_entry["content"])
assert content_json["result"]["id"] == "inc-123"
assert content_json["result"]["kpi_missed"][0]["kpi"] == "latency"
@pytest.mark.asyncio
async def test_otel_integration(
self,
callback_context,
):
"""Verifies OpenTelemetry integration in TraceManager."""
# Mock the tracer and span
mock_tracer = mock.Mock()
mock_span = mock.Mock()
mock_context = mock.Mock()
# Setup mock IDs (128-bit trace_id, 64-bit span_id)
trace_id_int = 0x12345678123456781234567812345678
span_id_int = 0x1234567812345678
mock_context.trace_id = trace_id_int
mock_context.span_id = span_id_int
mock_context.is_valid = True
mock_span.get_span_context.return_value = mock_context
mock_span.start_time = 1234567890000000000 # Mock start time in ns
mock_tracer.start_span.return_value = mock_span
# Patch the global tracer in the plugin module
with mock.patch(
"google.adk.plugins.bigquery_agent_analytics_plugin.tracer", mock_tracer
):
# Test push_span
span_id = bigquery_agent_analytics_plugin.TraceManager.push_span(
callback_context, "test_span"
)
mock_tracer.start_span.assert_called_with("test_span")
assert span_id == format(span_id_int, "016x")
# Test get_trace_id
# We need to mock trace.get_current_span() to return our mock span
# because push_span calls trace.attach(), which affects the global context
with mock.patch(
"opentelemetry.trace.get_current_span", return_value=mock_span
):
trace_id = bigquery_agent_analytics_plugin.TraceManager.get_trace_id(
callback_context
)
assert trace_id == format(trace_id_int, "032x")
# Test pop_span
# pop_span calls span.end()
bigquery_agent_analytics_plugin.TraceManager.pop_span()
mock_span.end.assert_called_once()
@pytest.mark.asyncio
async def test_otel_integration_real_provider(self, callback_context):
"""Verifies TraceManager with a real OpenTelemetry TracerProvider."""
# Setup OTEL with in-memory exporter
# pylint: disable=g-import-not-at-top
from opentelemetry.sdk import trace as trace_sdk
from opentelemetry.sdk.trace import export as trace_export
from opentelemetry.sdk.trace.export import in_memory_span_exporter
# pylint: enable=g-import-not-at-top
provider = trace_sdk.TracerProvider()
exporter = in_memory_span_exporter.InMemorySpanExporter()
processor = trace_export.SimpleSpanProcessor(exporter)
provider.add_span_processor(processor)
tracer = provider.get_tracer("test_tracer")
# Patch the global tracer in the plugin module
with mock.patch(
"google.adk.plugins.bigquery_agent_analytics_plugin.tracer", tracer
):
# 1. Start a span
span_id = bigquery_agent_analytics_plugin.TraceManager.push_span(
callback_context, "test_span"
)
# Verify a span was started but not ended
current_spans = exporter.get_finished_spans()
assert not current_spans
# Verify we can retrieve the trace ID
trace_id = bigquery_agent_analytics_plugin.TraceManager.get_trace_id(
callback_context
)
assert trace_id is not None
# 2. End the span
popped_span_id, _ = (
bigquery_agent_analytics_plugin.TraceManager.pop_span()
)
assert popped_span_id == span_id
# Verify span is now finished and exported
finished_spans = exporter.get_finished_spans()
assert len(finished_spans) == 1
assert finished_spans[0].name == "test_span"
assert format(finished_spans[0].context.span_id, "016x") == span_id
assert format(finished_spans[0].context.trace_id, "032x") == trace_id
@pytest.mark.asyncio
async def test_flush_mechanism(
self,
bq_plugin_inst,
mock_write_client,
dummy_arrow_schema,
invocation_context,
):
"""Verifies that flush() forces pending events to be written."""
# Log an event
bigquery_agent_analytics_plugin.TraceManager.push_span(invocation_context)
await bq_plugin_inst.before_run_callback(
invocation_context=invocation_context
)
# Call flush - this should block until the event is written
await bq_plugin_inst.flush()
# Verify write called
mock_write_client.append_rows.assert_called_once()
log_entry = await _get_captured_event_dict_async(
mock_write_client, dummy_arrow_schema
)
assert log_entry["event_type"] == "INVOCATION_STARTING"
@pytest.mark.asyncio
@pytest.mark.parametrize(
"gen_config_kwargs, expected_llm_config",
[
(
{
"temperature": 0.0,
"top_k": 5.0,
"top_p": 0.1,
"candidate_count": 5,
"max_output_tokens": 65000,
"presence_penalty": 0.1,
"frequency_penalty": 0.5,
"response_logprobs": True,
"logprobs": 3,
"seed": 42,
"labels": {"llm.agent.name": "test_agent"},
},
{
"temperature": 0.0,
"top_k": 5.0,
"top_p": 0.1,
"candidate_count": 5,
"max_output_tokens": 65000,
"presence_penalty": 0.1,
"frequency_penalty": 0.5,
"response_logprobs": True,
"logprobs": 3,
"seed": 42,
},
),
],
)
async def test_generation_config_logging(
self,
bq_plugin_inst,
mock_write_client,
dummy_arrow_schema,
callback_context,
gen_config_kwargs,
expected_llm_config,
):
"""Verifies that all fields in GenerateContentConfig are logged correctly."""
gen_config = types.GenerateContentConfig(**gen_config_kwargs)
llm_request = llm_request_lib.LlmRequest(
model="gemini-pro",
contents=[types.Content(parts=[types.Part(text="Prompt")])],
config=gen_config,
)
bigquery_agent_analytics_plugin.TraceManager.push_span(callback_context)
await bq_plugin_inst.before_model_callback(
callback_context=callback_context, llm_request=llm_request
)
# Flush
await bq_plugin_inst.flush()
# Verify
log_entry = await _get_captured_event_dict_async(
mock_write_client, dummy_arrow_schema
)
assert log_entry["event_type"] == "LLM_REQUEST"
attributes = json.loads(log_entry["attributes"])
llm_config = attributes.get("llm_config", {})
assert llm_config == expected_llm_config
if "labels" in gen_config_kwargs:
assert attributes.get("labels") == gen_config_kwargs["labels"]
class TestSafeCallbackDecorator:
"""Tests that _safe_callback prevents plugin errors from propagating."""
@pytest.mark.asyncio
async def test_callback_exception_does_not_propagate(
self,
bq_plugin_inst,
mock_write_client,
invocation_context,
):
"""A callback that throws should return None, not crash."""
# Force _log_event to raise
with mock.patch.object(
bq_plugin_inst,
"_log_event",
side_effect=RuntimeError("BQ network timeout"),
):
# Should NOT raise
result = await bq_plugin_inst.on_user_message_callback(
invocation_context=invocation_context,
user_message=types.Content(parts=[types.Part(text="Test")]),
)
assert result is None
@pytest.mark.asyncio
async def test_callback_exception_is_logged(
self,
bq_plugin_inst,
mock_write_client,
invocation_context,
):
"""The swallowed exception should be logged with exc_info."""
with mock.patch.object(
bq_plugin_inst,
"_log_event",
side_effect=RuntimeError("BQ write failed"),
):
with mock.patch(
"google.adk.plugins.bigquery_agent_analytics_plugin.logger"
) as mock_logger:
await bq_plugin_inst.before_run_callback(
invocation_context=invocation_context,
)
mock_logger.exception.assert_called_once_with(
"BigQuery analytics plugin error in %s; skipping.",
"before_run_callback",
)
@pytest.mark.asyncio
async def test_subsequent_callbacks_work_after_failure(
self,
bq_plugin_inst,
mock_write_client,
invocation_context,
dummy_arrow_schema,
):
"""After one callback fails, the next one should still work."""
call_count = 0
original_log_event = bq_plugin_inst._log_event
async def fail_once(*args, **kwargs):
nonlocal call_count
call_count += 1
if call_count == 1:
raise RuntimeError("Transient error")
return await original_log_event(*args, **kwargs)
with mock.patch.object(bq_plugin_inst, "_log_event", side_effect=fail_once):
# First call fails silently
await bq_plugin_inst.on_user_message_callback(
invocation_context=invocation_context,
user_message=types.Content(parts=[types.Part(text="Fail")]),
)
# Second call succeeds
bigquery_agent_analytics_plugin.TraceManager.push_span(invocation_context)
await bq_plugin_inst.before_run_callback(
invocation_context=invocation_context,
)
await asyncio.sleep(0.01)
mock_write_client.append_rows.assert_called_once()
@pytest.mark.asyncio
async def test_on_event_callback_exception_returns_none(
self,
bq_plugin_inst,
mock_write_client,
invocation_context,
):
"""on_event_callback should return None on error, not crash."""
event = event_lib.Event(
author="test_agent",
actions=event_actions_lib.EventActions(state_delta={"key": "value"}),
)
with mock.patch.object(
bq_plugin_inst,
"_log_event",
side_effect=Exception("serialize error"),
):
result = await bq_plugin_inst.on_event_callback(
invocation_context=invocation_context, event=event
)
assert result is None
@pytest.mark.asyncio
async def test_tool_callback_exception_does_not_propagate(
self,
bq_plugin_inst,
mock_write_client,
tool_context,
):
"""Tool callbacks should not crash even if plugin errors."""
mock_tool = mock.create_autospec(
base_tool_lib.BaseTool, instance=True, spec_set=True
)
type(mock_tool).name = mock.PropertyMock(return_value="MyTool")
with mock.patch.object(
bq_plugin_inst,
"_log_event",
side_effect=RuntimeError("BQ down"),
):
# before_tool_callback
result = await bq_plugin_inst.before_tool_callback(
tool=mock_tool,
tool_args={"p": "v"},
tool_context=tool_context,
)
assert result is None
# after_tool_callback
result = await bq_plugin_inst.after_tool_callback(
tool=mock_tool,
tool_args={"p": "v"},
tool_context=tool_context,
result={"r": "ok"},
)
assert result is None
# on_tool_error_callback
result = await bq_plugin_inst.on_tool_error_callback(
tool=mock_tool,
tool_args={"p": "v"},
tool_context=tool_context,
error=ValueError("tool broke"),
)
assert result is None
@pytest.mark.asyncio
async def test_model_callback_exception_does_not_propagate(
self,
bq_plugin_inst,
mock_write_client,
callback_context,
):
"""Model callbacks should not crash even if plugin errors."""
with mock.patch.object(
bq_plugin_inst,
"_log_event",
side_effect=RuntimeError("BQ down"),
):
llm_request = llm_request_lib.LlmRequest(
model="gemini-pro",
contents=[types.Content(parts=[types.Part(text="Hi")])],
)
result = await bq_plugin_inst.before_model_callback(
callback_context=callback_context, llm_request=llm_request
)
assert result is None
llm_response = llm_response_lib.LlmResponse(
content=types.Content(parts=[types.Part(text="Hi")]),
)
result = await bq_plugin_inst.after_model_callback(
callback_context=callback_context, llm_response=llm_response
)
assert result is None
result = await bq_plugin_inst.on_model_error_callback(
callback_context=callback_context,
llm_request=llm_request_lib.LlmRequest(model="gemini-pro"),
error=ValueError("llm error"),
)
assert result is None
class TestParserReuse:
"""Tests that HybridContentParser is reused, not recreated per event."""
@pytest.mark.asyncio
async def test_parser_instance_is_reused(
self,
bq_plugin_inst,
mock_write_client,
invocation_context,
):
"""The same parser instance should be reused across _log_event calls."""
parser_after_init = bq_plugin_inst.parser
assert parser_after_init is not None
bigquery_agent_analytics_plugin.TraceManager.push_span(invocation_context)
await bq_plugin_inst.on_user_message_callback(
invocation_context=invocation_context,
user_message=types.Content(parts=[types.Part(text="Hello")]),
)
await asyncio.sleep(0.01)
# Parser should be the same instance, not a new one
assert bq_plugin_inst.parser is parser_after_init
@pytest.mark.asyncio
async def test_parser_trace_id_updated_per_call(
self,
bq_plugin_inst,
mock_write_client,
invocation_context,
dummy_arrow_schema,
):
"""trace_id and span_id on the parser should update per _log_event."""
parser = bq_plugin_inst.parser
original_trace_id = parser.trace_id
bigquery_agent_analytics_plugin.TraceManager.push_span(invocation_context)
await bq_plugin_inst.on_user_message_callback(
invocation_context=invocation_context,
user_message=types.Content(parts=[types.Part(text="Test")]),
)
await asyncio.sleep(0.01)
# After logging, trace_id/span_id should have been updated
# (they're derived from TraceManager, not the initial empty strings)
assert parser.span_id != ""
@pytest.mark.asyncio
async def test_parser_not_recreated_with_constructor(
self,
bq_plugin_inst,
mock_write_client,
invocation_context,
):
"""HybridContentParser constructor should not be called in
_log_event."""
with mock.patch.object(
bigquery_agent_analytics_plugin,
"HybridContentParser",
wraps=bigquery_agent_analytics_plugin.HybridContentParser,
) as mock_parser_cls:
bigquery_agent_analytics_plugin.TraceManager.push_span(invocation_context)
await bq_plugin_inst.on_user_message_callback(
invocation_context=invocation_context,
user_message=types.Content(parts=[types.Part(text="Test")]),
)
await asyncio.sleep(0.01)
# Constructor should NOT have been called during _log_event
mock_parser_cls.assert_not_called()
class TestPropertyAccessors:
"""Tests that properties work correctly after __getattribute__ removal."""
@pytest.mark.asyncio
async def testbatch_processorerty_returns_processor(self, bq_plugin_inst):
"""batch_processor property should return the processor for the
current loop."""
bp = bq_plugin_inst.batch_processor
assert bp is not None
assert isinstance(bp, bigquery_agent_analytics_plugin.BatchProcessor)
@pytest.mark.asyncio
async def test_write_client_property_returns_client(self, bq_plugin_inst):
"""write_client property should return the client for the current
loop."""
wc = bq_plugin_inst.write_client
assert wc is not None
@pytest.mark.asyncio
async def test_write_stream_property_returns_stream(self, bq_plugin_inst):
"""write_stream property should return the stream name."""
ws = bq_plugin_inst.write_stream
assert ws is not None
assert ws == DEFAULT_STREAM_NAME
@pytest.mark.asyncio
async def test_properties_return_none_when_no_loop_state(self):
"""Properties should return None when no state exists for the
current loop."""
plugin = bigquery_agent_analytics_plugin.BigQueryAgentAnalyticsPlugin(
project_id=PROJECT_ID,
dataset_id=DATASET_ID,
table_id=TABLE_ID,
)
assert plugin.batch_processor is None
assert plugin.write_client is None
assert plugin.write_stream is None
@pytest.mark.asyncio
async def test_regular_attributes_still_accessible(self, bq_plugin_inst):
"""Regular instance attributes should still be accessible."""
assert bq_plugin_inst.project_id == PROJECT_ID
assert bq_plugin_inst.dataset_id == DATASET_ID
assert bq_plugin_inst.table_id == TABLE_ID
assert bq_plugin_inst.config is not None
assert bq_plugin_inst._started is True
def test_properties_without_running_loop(self):
"""Properties should return None when no event loop is running."""
plugin = bigquery_agent_analytics_plugin.BigQueryAgentAnalyticsPlugin(
project_id=PROJECT_ID,
dataset_id=DATASET_ID,
table_id=TABLE_ID,
)
# No running loop → should return None, not crash
assert plugin.batch_processor is None
assert plugin.write_client is None
assert plugin.write_stream is None
class TestUnifiedSpanRecords:
"""Tests for the unified _SpanRecord-based TraceManager."""
@pytest.mark.asyncio
async def test_push_pop_keeps_stacks_in_sync(self, callback_context):
"""Push and pop should always leave the records stack consistent."""
TM = bigquery_agent_analytics_plugin.TraceManager
TM.init_trace(callback_context)
span_id_1 = TM.push_span(callback_context, "span-1")
span_id_2 = TM.push_span(callback_context, "span-2")
# Both should be on the stack
assert TM.get_current_span_id() == span_id_2
current, parent = TM.get_current_span_and_parent()
assert current == span_id_2
assert parent == span_id_1
# Pop span-2
popped_id, duration = TM.pop_span()
assert popped_id == span_id_2
assert duration is not None
assert TM.get_current_span_id() == span_id_1
# Pop span-1
popped_id, _ = TM.pop_span()
assert popped_id == span_id_1
assert TM.get_current_span_id() is None
@pytest.mark.asyncio
async def test_pop_empty_stack_returns_none(self, callback_context):
"""Popping an empty stack should return (None, None)."""
TM = bigquery_agent_analytics_plugin.TraceManager
TM.init_trace(callback_context)
span_id, duration = TM.pop_span()
assert span_id is None
assert duration is None
@pytest.mark.asyncio
async def test_first_token_time_stored_in_record(self, callback_context):
"""first_token_time should be stored on the span record."""
TM = bigquery_agent_analytics_plugin.TraceManager
TM.init_trace(callback_context)
span_id = TM.push_span(callback_context, "llm-span")
# No first token yet
assert TM.get_first_token_time(span_id) is None
# Record first token
assert TM.record_first_token(span_id) is True
ftt = TM.get_first_token_time(span_id)
assert ftt is not None
# Second call should return False (already recorded)
assert TM.record_first_token(span_id) is False
# Clean up
TM.pop_span()
@pytest.mark.asyncio
async def test_start_time_accessible_by_span_id(self, callback_context):
"""get_start_time should find the span by ID in the records."""
TM = bigquery_agent_analytics_plugin.TraceManager
TM.init_trace(callback_context)
span_id = TM.push_span(callback_context, "timed-span")
start = TM.get_start_time(span_id)
assert start is not None
assert start > 0
TM.pop_span()
@pytest.mark.asyncio
async def test_attach_current_span_does_not_own(self, callback_context):
"""attach_current_span should not end the span on pop."""
TM = bigquery_agent_analytics_plugin.TraceManager
TM.init_trace(callback_context)
mock_span = mock.Mock()
mock_ctx = mock.Mock()
mock_ctx.is_valid = False
mock_span.get_span_context.return_value = mock_ctx
with mock.patch(
"opentelemetry.trace.get_current_span", return_value=mock_span
):
span_id = TM.attach_current_span(callback_context)
assert span_id is not None
TM.pop_span()
# Should NOT have called span.end() since we don't own it
mock_span.end.assert_not_called()
@pytest.mark.asyncio
async def test_concurrent_tasks_have_isolated_stacks(self, callback_context):
"""Concurrent async tasks should have isolated span stacks."""
TM = bigquery_agent_analytics_plugin.TraceManager
TM.init_trace(callback_context)
async def task_a():
s = TM.push_span(callback_context, "task-a")
await asyncio.sleep(0.02)
assert TM.get_current_span_id() == s
TM.pop_span()
return s
async def task_b():
s = TM.push_span(callback_context, "task-b")
await asyncio.sleep(0.02)
assert TM.get_current_span_id() == s
TM.pop_span()
return s
results = await asyncio.gather(task_a(), task_b())
assert results[0] != results[1]
@pytest.mark.asyncio
async def test_pop_cleans_up_record_completely(self, callback_context):
"""After pop, the record should be fully removed from the stack."""
TM = bigquery_agent_analytics_plugin.TraceManager
TM.init_trace(callback_context)
span_id = TM.push_span(callback_context, "temp-span")
# Record is on the stack
assert TM.get_current_span_id() == span_id
assert TM.get_start_time(span_id) is not None
TM.pop_span()
# Record is gone
assert TM.get_current_span_id() is None
assert TM.get_start_time(span_id) is None
assert TM.get_first_token_time(span_id) is None
class TestLoopStateValidation:
"""Tests for loop state validation and stale loop cleanup."""
def _make_plugin(self):
"""Creates a plugin instance without starting it."""
return bigquery_agent_analytics_plugin.BigQueryAgentAnalyticsPlugin(
project_id=PROJECT_ID,
dataset_id=DATASET_ID,
table_id=TABLE_ID,
)
def _make_loop_state(self):
"""Creates a mock _LoopState with batch_processor and write_client."""
state = mock.MagicMock()
state.batch_processor = mock.MagicMock(
spec=bigquery_agent_analytics_plugin.BatchProcessor
)
state.batch_processor.flush = mock.AsyncMock()
state.write_client = mock.MagicMock()
return state
def test_cleanup_stale_loop_states_removes_closed_loops(self):
"""Closed loops should be removed from _loop_state_by_loop."""
plugin = self._make_plugin()
closed_loop = mock.MagicMock(spec=asyncio.AbstractEventLoop)
closed_loop.is_closed.return_value = True
plugin._loop_state_by_loop[closed_loop] = self._make_loop_state()
plugin._cleanup_stale_loop_states()
assert closed_loop not in plugin._loop_state_by_loop
def test_cleanup_stale_loop_states_keeps_open_loops(self):
"""Open loops should not be removed from _loop_state_by_loop."""
plugin = self._make_plugin()
open_loop = mock.MagicMock(spec=asyncio.AbstractEventLoop)
open_loop.is_closed.return_value = False
plugin._loop_state_by_loop[open_loop] = self._make_loop_state()
plugin._cleanup_stale_loop_states()
assert open_loop in plugin._loop_state_by_loop
def test_cleanup_removes_only_closed_loops(self):
"""Only closed loops should be removed; open ones stay."""
plugin = self._make_plugin()
open_loop = mock.MagicMock(spec=asyncio.AbstractEventLoop)
open_loop.is_closed.return_value = False
closed_loop = mock.MagicMock(spec=asyncio.AbstractEventLoop)
closed_loop.is_closed.return_value = True
plugin._loop_state_by_loop[open_loop] = self._make_loop_state()
plugin._loop_state_by_loop[closed_loop] = self._make_loop_state()
plugin._cleanup_stale_loop_states()
assert open_loop in plugin._loop_state_by_loop
assert closed_loop not in plugin._loop_state_by_loop
@pytest.mark.asyncio
async def testbatch_processor_returns_processor_for_open_loop(
self,
):
"""batch_processor returns processor for the current loop."""
plugin = self._make_plugin()
loop = asyncio.get_running_loop()
state = self._make_loop_state()
plugin._loop_state_by_loop[loop] = state
assert plugin.batch_processor is state.batch_processor
# Clean up
del plugin._loop_state_by_loop[loop]
@pytest.mark.asyncio
async def testbatch_processor_cleans_closed_loop_entry(self):
"""Accessing batch_processor cleans up closed loop entries."""
plugin = self._make_plugin()
closed_loop = mock.MagicMock(spec=asyncio.AbstractEventLoop)
closed_loop.is_closed.return_value = True
plugin._loop_state_by_loop[closed_loop] = self._make_loop_state()
# Accessing the prop should clean up the closed loop entry
_ = plugin.batch_processor
assert closed_loop not in plugin._loop_state_by_loop
@pytest.mark.asyncio
async def test_flush_cleans_stale_states(self):
"""flush() should clean up stale loop states before flushing."""
plugin = self._make_plugin()
closed_loop = mock.MagicMock(spec=asyncio.AbstractEventLoop)
closed_loop.is_closed.return_value = True
plugin._loop_state_by_loop[closed_loop] = self._make_loop_state()
await plugin.flush()
assert closed_loop not in plugin._loop_state_by_loop
class TestAtexitCleanup:
"""Tests for the simplified _atexit_cleanup static method."""
def _make_batch_processor(self, queue_items=0):
bp = mock.MagicMock()
bp._shutdown = False
q = asyncio.Queue()
for i in range(queue_items):
q.put_nowait({"event": i})
bp._queue = q
return bp
def test_skips_none_processor(self):
"""Should return immediately when batch_processor is None."""
bigquery_agent_analytics_plugin.BigQueryAgentAnalyticsPlugin._atexit_cleanup(
None
)
def test_skips_already_shutdown(self):
"""Should return immediately when batch_processor._shutdown is True."""
bp = self._make_batch_processor()
bp._shutdown = True
bigquery_agent_analytics_plugin.BigQueryAgentAnalyticsPlugin._atexit_cleanup(
bp
)
def test_skips_reference_error(self):
"""Should handle ReferenceError from weakref'd processor."""
bp = mock.MagicMock()
type(bp)._shutdown = mock.PropertyMock(side_effect=ReferenceError)
bigquery_agent_analytics_plugin.BigQueryAgentAnalyticsPlugin._atexit_cleanup(
bp
)
def test_empty_queue_no_warning(self):
"""Should not warn when queue is empty."""
bp = self._make_batch_processor(queue_items=0)
with mock.patch.object(
bigquery_agent_analytics_plugin.logger, "warning"
) as mock_warn:
bigquery_agent_analytics_plugin.BigQueryAgentAnalyticsPlugin._atexit_cleanup(
bp
)
mock_warn.assert_not_called()
def test_remaining_items_logs_warning(self):
"""Should drain queue and log warning with count of lost items."""
bp = self._make_batch_processor(queue_items=3)
with mock.patch.object(
bigquery_agent_analytics_plugin.logger, "warning"
) as mock_warn:
bigquery_agent_analytics_plugin.BigQueryAgentAnalyticsPlugin._atexit_cleanup(
bp
)
mock_warn.assert_called_once()
# Verify the warning mentions the count
call_args = mock_warn.call_args
assert "3" in str(call_args)
def test_queue_is_drained(self):
"""Should drain all items from the queue."""
bp = self._make_batch_processor(queue_items=5)
bigquery_agent_analytics_plugin.BigQueryAgentAnalyticsPlugin._atexit_cleanup(
bp
)
assert bp._queue.empty()
class TestDuplicateLabels:
"""Tests that labels in before_model_callback are set exactly once."""
@pytest.mark.asyncio
async def test_labels_set_when_present(
self,
bq_plugin_inst,
mock_write_client,
callback_context,
dummy_arrow_schema,
):
"""Labels should appear in attributes when config has them."""
llm_request = llm_request_lib.LlmRequest(
model="gemini-pro",
config=types.GenerateContentConfig(
labels={"env": "test"},
),
contents=[types.Content(role="user", parts=[types.Part(text="hi")])],
)
bigquery_agent_analytics_plugin.TraceManager.push_span(callback_context)
await bq_plugin_inst.before_model_callback(
callback_context=callback_context, llm_request=llm_request
)
await asyncio.sleep(0.01)
log_entry = await _get_captured_event_dict_async(
mock_write_client, dummy_arrow_schema
)
attributes = json.loads(log_entry["attributes"])
assert attributes["labels"] == {"env": "test"}
@pytest.mark.asyncio
async def test_labels_absent_when_none(
self,
bq_plugin_inst,
mock_write_client,
callback_context,
dummy_arrow_schema,
):
"""Labels should not appear in attributes when config.labels is None."""
llm_request = llm_request_lib.LlmRequest(
model="gemini-pro",
config=types.GenerateContentConfig(
temperature=0.5,
),
contents=[types.Content(role="user", parts=[types.Part(text="hi")])],
)
bigquery_agent_analytics_plugin.TraceManager.push_span(callback_context)
await bq_plugin_inst.before_model_callback(
callback_context=callback_context, llm_request=llm_request
)
await asyncio.sleep(0.01)
log_entry = await _get_captured_event_dict_async(
mock_write_client, dummy_arrow_schema
)
attributes = json.loads(log_entry["attributes"])
assert "labels" not in attributes
@pytest.mark.asyncio
async def test_no_config_no_labels(
self,
bq_plugin_inst,
mock_write_client,
callback_context,
dummy_arrow_schema,
):
"""Labels should not appear when llm_request has no config."""
llm_request = llm_request_lib.LlmRequest(
model="gemini-pro",
contents=[types.Content(role="user", parts=[types.Part(text="hi")])],
)
bigquery_agent_analytics_plugin.TraceManager.push_span(callback_context)
await bq_plugin_inst.before_model_callback(
callback_context=callback_context, llm_request=llm_request
)
await asyncio.sleep(0.01)
log_entry = await _get_captured_event_dict_async(
mock_write_client, dummy_arrow_schema
)
attributes = json.loads(log_entry["attributes"])
assert "labels" not in attributes
class TestResolveSpanIds:
"""Tests for the _resolve_span_ids static helper."""
def test_uses_trace_manager_defaults(self):
"""Should use TraceManager values when no overrides provided."""
ed = bigquery_agent_analytics_plugin.EventData(
extra_attributes={"some_key": "value"}
)
with mock.patch.object(
bigquery_agent_analytics_plugin.TraceManager,
"get_current_span_and_parent",
return_value=("span-1", "parent-1"),
):
span_id, parent_id = (
bigquery_agent_analytics_plugin.BigQueryAgentAnalyticsPlugin._resolve_span_ids(
ed
)
)
assert span_id == "span-1"
assert parent_id == "parent-1"
def test_span_id_override(self):
"""Should use span_id_override from EventData."""
ed = bigquery_agent_analytics_plugin.EventData(
span_id_override="custom-span"
)
with mock.patch.object(
bigquery_agent_analytics_plugin.TraceManager,
"get_current_span_and_parent",
return_value=("span-1", "parent-1"),
):
span_id, parent_id = (
bigquery_agent_analytics_plugin.BigQueryAgentAnalyticsPlugin._resolve_span_ids(
ed
)
)
assert span_id == "custom-span"
assert parent_id == "parent-1"
def test_parent_span_id_override(self):
"""Should use parent_span_id_override from EventData."""
ed = bigquery_agent_analytics_plugin.EventData(
parent_span_id_override="custom-parent"
)
with mock.patch.object(
bigquery_agent_analytics_plugin.TraceManager,
"get_current_span_and_parent",
return_value=("span-1", "parent-1"),
):
span_id, parent_id = (
bigquery_agent_analytics_plugin.BigQueryAgentAnalyticsPlugin._resolve_span_ids(
ed
)
)
assert span_id == "span-1"
assert parent_id == "custom-parent"
def test_none_override_keeps_default(self):
"""None overrides should keep the TraceManager defaults."""
ed = bigquery_agent_analytics_plugin.EventData(
span_id_override=None, parent_span_id_override=None
)
with mock.patch.object(
bigquery_agent_analytics_plugin.TraceManager,
"get_current_span_and_parent",
return_value=("span-1", "parent-1"),
):
span_id, parent_id = (
bigquery_agent_analytics_plugin.BigQueryAgentAnalyticsPlugin._resolve_span_ids(
ed
)
)
assert span_id == "span-1"
assert parent_id == "parent-1"
class TestExtractLatency:
"""Tests for the _extract_latency static helper."""
def test_no_latency_returns_none(self):
"""Should return None when no latency fields present."""
ed = bigquery_agent_analytics_plugin.EventData(
extra_attributes={"other": "val"}
)
result = bigquery_agent_analytics_plugin.BigQueryAgentAnalyticsPlugin._extract_latency(
ed
)
assert result is None
def test_total_latency_only(self):
"""Should extract latency_ms into total_ms."""
ed = bigquery_agent_analytics_plugin.EventData(latency_ms=42.5)
result = bigquery_agent_analytics_plugin.BigQueryAgentAnalyticsPlugin._extract_latency(
ed
)
assert result == {"total_ms": 42.5}
def test_tfft_only(self):
"""Should extract time_to_first_token_ms."""
ed = bigquery_agent_analytics_plugin.EventData(time_to_first_token_ms=10.0)
result = bigquery_agent_analytics_plugin.BigQueryAgentAnalyticsPlugin._extract_latency(
ed
)
assert result == {"time_to_first_token_ms": 10.0}
def test_both_latencies(self):
"""Should extract both latency fields."""
ed = bigquery_agent_analytics_plugin.EventData(
latency_ms=100, time_to_first_token_ms=20
)
result = bigquery_agent_analytics_plugin.BigQueryAgentAnalyticsPlugin._extract_latency(
ed
)
assert result == {"total_ms": 100, "time_to_first_token_ms": 20}
class TestEnrichAttributes:
"""Tests for the _enrich_attributes helper."""
def _make_plugin(self):
with (
mock.patch(
"google.auth.default",
return_value=(mock.Mock(), PROJECT_ID),
),
mock.patch(
"google.cloud.bigquery.Client",
),
):
plugin = bigquery_agent_analytics_plugin.BigQueryAgentAnalyticsPlugin(
project_id=PROJECT_ID,
dataset_id=DATASET_ID,
)
plugin.config.max_content_length = 10000
plugin.config.log_session_metadata = False
plugin.config.custom_tags = None
return plugin
def _make_callback_context(self):
ctx = mock.MagicMock()
session = mock.MagicMock()
session.id = "sess-001"
session.app_name = "test-app"
session.user_id = "user-001"
session.state = {"env": "test"}
ctx._invocation_context.session = session
return ctx
def test_adds_root_agent_name(self):
"""Should always add root_agent_name."""
plugin = self._make_plugin()
ed = bigquery_agent_analytics_plugin.EventData()
with mock.patch.object(
bigquery_agent_analytics_plugin.TraceManager,
"get_root_agent_name",
return_value="my-agent",
):
attrs = plugin._enrich_attributes(ed, self._make_callback_context())
assert attrs["root_agent_name"] == "my-agent"
def test_includes_model(self):
"""Should include model from EventData."""
plugin = self._make_plugin()
ed = bigquery_agent_analytics_plugin.EventData(model="gemini-pro")
with mock.patch.object(
bigquery_agent_analytics_plugin.TraceManager,
"get_root_agent_name",
return_value="agent",
):
attrs = plugin._enrich_attributes(ed, self._make_callback_context())
assert attrs["model"] == "gemini-pro"
def test_session_metadata_when_enabled(self):
"""Should add session_metadata when log_session_metadata is True."""
plugin = self._make_plugin()
plugin.config.log_session_metadata = True
ctx = self._make_callback_context()
ed = bigquery_agent_analytics_plugin.EventData()
with mock.patch.object(
bigquery_agent_analytics_plugin.TraceManager,
"get_root_agent_name",
return_value="agent",
):
attrs = plugin._enrich_attributes(ed, ctx)
meta = attrs["session_metadata"]
assert meta["session_id"] == "sess-001"
assert meta["app_name"] == "test-app"
assert meta["user_id"] == "user-001"
assert meta["state"] == {"env": "test"}
def test_session_metadata_when_disabled(self):
"""Should not add session_metadata when log_session_metadata is False."""
plugin = self._make_plugin()
plugin.config.log_session_metadata = False
ed = bigquery_agent_analytics_plugin.EventData()
with mock.patch.object(
bigquery_agent_analytics_plugin.TraceManager,
"get_root_agent_name",
return_value="agent",
):
attrs = plugin._enrich_attributes(ed, self._make_callback_context())
assert "session_metadata" not in attrs
def test_custom_tags_added(self):
"""Should add custom_tags when configured."""
plugin = self._make_plugin()
plugin.config.custom_tags = {"team": "infra"}
ed = bigquery_agent_analytics_plugin.EventData()
with mock.patch.object(
bigquery_agent_analytics_plugin.TraceManager,
"get_root_agent_name",
return_value="agent",
):
attrs = plugin._enrich_attributes(ed, self._make_callback_context())
assert attrs["custom_tags"] == {"team": "infra"}
def test_usage_metadata_truncated(self):
"""Should smart-truncate usage_metadata."""
plugin = self._make_plugin()
ed = bigquery_agent_analytics_plugin.EventData(
usage_metadata={"input_tokens": 100, "output_tokens": 50}
)
with mock.patch.object(
bigquery_agent_analytics_plugin.TraceManager,
"get_root_agent_name",
return_value="agent",
):
attrs = plugin._enrich_attributes(ed, self._make_callback_context())
assert attrs["usage_metadata"] == {
"input_tokens": 100,
"output_tokens": 50,
}
class TestMultiSubagentToolLogging:
"""Tests that tool events from different subagents are attributed correctly.
Covers:
- Tool calls from different subagents have the correct `agent` field
- Multi-turn (different invocation_ids, same session) logs correctly
- Full callback sequence across multiple subagents in one turn
- Span hierarchy is maintained per-subagent
"""
@staticmethod
def _make_invocation_context(agent_name, session, invocation_id="inv-001"):
"""Create an InvocationContext with a specific agent name."""
mock_a = mock.create_autospec(
base_agent.BaseAgent, instance=True, spec_set=True
)
type(mock_a).name = mock.PropertyMock(return_value=agent_name)
type(mock_a).instruction = mock.PropertyMock(
return_value=f"{agent_name} instruction"
)
mock_session_service = mock.create_autospec(
base_session_service_lib.BaseSessionService,
instance=True,
spec_set=True,
)
mock_plugin_manager = mock.create_autospec(
plugin_manager_lib.PluginManager,
instance=True,
spec_set=True,
)
return invocation_context_lib.InvocationContext(
agent=mock_a,
session=session,
invocation_id=invocation_id,
session_service=mock_session_service,
plugin_manager=mock_plugin_manager,
)
@staticmethod
def _make_session(session_id="session-multi", user_id="user-multi"):
mock_s = mock.create_autospec(
session_lib.Session, instance=True, spec_set=True
)
type(mock_s).id = mock.PropertyMock(return_value=session_id)
type(mock_s).user_id = mock.PropertyMock(return_value=user_id)
type(mock_s).app_name = mock.PropertyMock(return_value="test_app")
type(mock_s).state = mock.PropertyMock(return_value={})
return mock_s
@staticmethod
def _make_tool(name):
mock_tool = mock.create_autospec(
base_tool_lib.BaseTool, instance=True, spec_set=True
)
type(mock_tool).name = mock.PropertyMock(return_value=name)
type(mock_tool).description = mock.PropertyMock(
return_value=f"{name} description"
)
return mock_tool
@pytest.mark.asyncio
async def test_tool_calls_attributed_to_correct_subagent(
self,
mock_auth_default,
mock_bq_client,
mock_write_client,
mock_to_arrow_schema,
dummy_arrow_schema,
mock_asyncio_to_thread,
):
"""Tool events from different subagents carry the correct agent name."""
session = self._make_session()
async with managed_plugin(
PROJECT_ID, DATASET_ID, table_id=TABLE_ID
) as plugin:
await plugin._ensure_started()
mock_write_client.append_rows.reset_mock()
# --- Subagent A: schema_explorer calls list_datasets ---
inv_ctx_a = self._make_invocation_context("schema_explorer", session)
ctx_a = tool_context_lib.ToolContext(invocation_context=inv_ctx_a)
tool_a = self._make_tool("list_dataset_ids")
bigquery_agent_analytics_plugin.TraceManager.push_span(ctx_a, "tool")
await plugin.before_tool_callback(
tool=tool_a,
tool_args={"project_id": "my-project"},
tool_context=ctx_a,
)
await asyncio.sleep(0.01)
# --- Subagent B: image_describer calls describe_this_image ---
inv_ctx_b = self._make_invocation_context("image_describer", session)
ctx_b = tool_context_lib.ToolContext(invocation_context=inv_ctx_b)
tool_b = self._make_tool("describe_this_image")
bigquery_agent_analytics_plugin.TraceManager.push_span(ctx_b, "tool")
await plugin.before_tool_callback(
tool=tool_b,
tool_args={"image_uri": "gs://bucket/image.jpg"},
tool_context=ctx_b,
)
await asyncio.sleep(0.01)
rows = await _get_captured_rows_async(
mock_write_client, dummy_arrow_schema
)
assert len(rows) == 2
# First row: schema_explorer's tool
assert rows[0]["event_type"] == "TOOL_STARTING"
assert rows[0]["agent"] == "schema_explorer"
content_a = json.loads(rows[0]["content"])
assert content_a["tool"] == "list_dataset_ids"
assert content_a["args"] == {"project_id": "my-project"}
# Second row: image_describer's tool
assert rows[1]["event_type"] == "TOOL_STARTING"
assert rows[1]["agent"] == "image_describer"
content_b = json.loads(rows[1]["content"])
assert content_b["tool"] == "describe_this_image"
assert content_b["args"] == {"image_uri": "gs://bucket/image.jpg"}
@pytest.mark.asyncio
async def test_multi_turn_tool_calls_different_invocations(
self,
mock_auth_default,
mock_bq_client,
mock_write_client,
mock_to_arrow_schema,
dummy_arrow_schema,
mock_asyncio_to_thread,
):
"""Multi-turn: same session, different invocation IDs, tools logged."""
session = self._make_session()
async with managed_plugin(
PROJECT_ID, DATASET_ID, table_id=TABLE_ID
) as plugin:
await plugin._ensure_started()
mock_write_client.append_rows.reset_mock()
# --- Turn 1: schema_explorer calls list_dataset_ids ---
inv_ctx_1 = self._make_invocation_context(
"schema_explorer", session, invocation_id="inv-turn1"
)
ctx_1 = tool_context_lib.ToolContext(invocation_context=inv_ctx_1)
tool_1 = self._make_tool("list_dataset_ids")
bigquery_agent_analytics_plugin.TraceManager.push_span(ctx_1, "tool")
await plugin.before_tool_callback(
tool=tool_1,
tool_args={"project_id": "proj"},
tool_context=ctx_1,
)
await asyncio.sleep(0.01)
await plugin.after_tool_callback(
tool=tool_1,
tool_args={"project_id": "proj"},
tool_context=ctx_1,
result={"datasets": ["ds1", "ds2"]},
)
await asyncio.sleep(0.01)
# --- Turn 2: query_analyst calls execute_sql ---
inv_ctx_2 = self._make_invocation_context(
"query_analyst", session, invocation_id="inv-turn2"
)
ctx_2 = tool_context_lib.ToolContext(invocation_context=inv_ctx_2)
tool_2 = self._make_tool("execute_sql")
bigquery_agent_analytics_plugin.TraceManager.push_span(ctx_2, "tool")
await plugin.before_tool_callback(
tool=tool_2,
tool_args={"sql": "SELECT * FROM t"},
tool_context=ctx_2,
)
await asyncio.sleep(0.01)
await plugin.after_tool_callback(
tool=tool_2,
tool_args={"sql": "SELECT * FROM t"},
tool_context=ctx_2,
result={"rows": [{"col": "val"}]},
)
await asyncio.sleep(0.01)
rows = await _get_captured_rows_async(
mock_write_client, dummy_arrow_schema
)
assert len(rows) == 4
# Turn 1: TOOL_STARTING + TOOL_COMPLETED for schema_explorer
assert rows[0]["event_type"] == "TOOL_STARTING"
assert rows[0]["agent"] == "schema_explorer"
assert rows[0]["invocation_id"] == "inv-turn1"
assert rows[0]["session_id"] == "session-multi"
assert rows[1]["event_type"] == "TOOL_COMPLETED"
assert rows[1]["agent"] == "schema_explorer"
assert rows[1]["invocation_id"] == "inv-turn1"
content_1 = json.loads(rows[1]["content"])
assert content_1["tool"] == "list_dataset_ids"
assert content_1["result"] == {"datasets": ["ds1", "ds2"]}
# Turn 2: TOOL_STARTING + TOOL_COMPLETED for query_analyst
assert rows[2]["event_type"] == "TOOL_STARTING"
assert rows[2]["agent"] == "query_analyst"
assert rows[2]["invocation_id"] == "inv-turn2"
assert rows[3]["event_type"] == "TOOL_COMPLETED"
assert rows[3]["agent"] == "query_analyst"
assert rows[3]["invocation_id"] == "inv-turn2"
content_2 = json.loads(rows[3]["content"])
assert content_2["tool"] == "execute_sql"
assert content_2["result"] == {"rows": [{"col": "val"}]}
@pytest.mark.asyncio
async def test_full_subagent_callback_sequence(
self,
mock_auth_default,
mock_bq_client,
mock_write_client,
mock_to_arrow_schema,
dummy_arrow_schema,
mock_asyncio_to_thread,
):
"""Full lifecycle: agent_start → LLM → tool → tool_done → LLM → agent_done.
Simulates a subagent that makes an LLM call, then a tool call,
then another LLM call, and completes.
"""
session = self._make_session()
inv_ctx = self._make_invocation_context("schema_explorer", session)
cb_ctx = callback_context_lib.CallbackContext(invocation_context=inv_ctx)
tool_ctx = tool_context_lib.ToolContext(invocation_context=inv_ctx)
mock_agent = inv_ctx.agent
tool = self._make_tool("get_table_info")
async with managed_plugin(
PROJECT_ID, DATASET_ID, table_id=TABLE_ID
) as plugin:
await plugin._ensure_started()
mock_write_client.append_rows.reset_mock()
# 1. AGENT_STARTING
await plugin.before_agent_callback(
agent=mock_agent, callback_context=cb_ctx
)
await asyncio.sleep(0.01)
# 2. LLM_REQUEST (agent decides to call a tool)
llm_req = llm_request_lib.LlmRequest(
model="gemini-2.5-flash",
contents=[
types.Content(parts=[types.Part(text="What tables exist?")])
],
)
await plugin.before_model_callback(
callback_context=cb_ctx, llm_request=llm_req
)
await asyncio.sleep(0.01)
# 3. LLM_RESPONSE (function call)
llm_resp = llm_response_lib.LlmResponse(
content=types.Content(
parts=[
types.Part(
function_call=types.FunctionCall(
name="get_table_info",
args={"table": "events"},
)
)
]
)
)
await plugin.after_model_callback(
callback_context=cb_ctx, llm_response=llm_resp
)
await asyncio.sleep(0.01)
# 4. TOOL_STARTING
bigquery_agent_analytics_plugin.TraceManager.push_span(tool_ctx, "tool")
await plugin.before_tool_callback(
tool=tool,
tool_args={"table": "events"},
tool_context=tool_ctx,
)
await asyncio.sleep(0.01)
# 5. TOOL_COMPLETED
await plugin.after_tool_callback(
tool=tool,
tool_args={"table": "events"},
tool_context=tool_ctx,
result={"schema": [{"name": "id", "type": "INT64"}]},
)
await asyncio.sleep(0.01)
# 6. AGENT_COMPLETED
await plugin.after_agent_callback(
agent=mock_agent, callback_context=cb_ctx
)
await asyncio.sleep(0.01)
rows = await _get_captured_rows_async(
mock_write_client, dummy_arrow_schema
)
assert len(rows) == 6
expected_sequence = [
"AGENT_STARTING",
"LLM_REQUEST",
"LLM_RESPONSE",
"TOOL_STARTING",
"TOOL_COMPLETED",
"AGENT_COMPLETED",
]
for i, expected_type in enumerate(expected_sequence):
assert (
rows[i]["event_type"] == expected_type
), f"Row {i}: expected {expected_type}, got {rows[i]['event_type']}"
assert rows[i]["agent"] == "schema_explorer"
assert rows[i]["session_id"] == "session-multi"
# TOOL rows have correct content
tool_start = json.loads(rows[3]["content"])
assert tool_start["tool"] == "get_table_info"
assert tool_start["args"] == {"table": "events"}
tool_done = json.loads(rows[4]["content"])
assert tool_done["tool"] == "get_table_info"
assert tool_done["result"] == {"schema": [{"name": "id", "type": "INT64"}]}
# AGENT_COMPLETED and TOOL_COMPLETED should have latency
assert rows[4]["latency_ms"] is not None # TOOL_COMPLETED
assert rows[5]["latency_ms"] is not None # AGENT_COMPLETED
@pytest.mark.asyncio
async def test_tool_error_attributed_to_subagent(
self,
mock_auth_default,
mock_bq_client,
mock_write_client,
mock_to_arrow_schema,
dummy_arrow_schema,
mock_asyncio_to_thread,
):
"""TOOL_ERROR events carry the correct subagent name."""
session = self._make_session()
inv_ctx = self._make_invocation_context("query_analyst", session)
tool_ctx = tool_context_lib.ToolContext(invocation_context=inv_ctx)
tool = self._make_tool("execute_sql")
async with managed_plugin(
PROJECT_ID, DATASET_ID, table_id=TABLE_ID
) as plugin:
await plugin._ensure_started()
mock_write_client.append_rows.reset_mock()
bigquery_agent_analytics_plugin.TraceManager.push_span(tool_ctx, "tool")
await plugin.on_tool_error_callback(
tool=tool,
tool_args={"sql": "SELECT * FROM bad_table"},
tool_context=tool_ctx,
error=RuntimeError("Table not found"),
)
await asyncio.sleep(0.01)
rows = await _get_captured_rows_async(
mock_write_client, dummy_arrow_schema
)
assert len(rows) == 1
assert rows[0]["event_type"] == "TOOL_ERROR"
assert rows[0]["agent"] == "query_analyst"
assert rows[0]["error_message"] == "Table not found"
content = json.loads(rows[0]["content"])
assert content["tool"] == "execute_sql"
assert content["args"] == {"sql": "SELECT * FROM bad_table"}
@pytest.mark.asyncio
async def test_multi_subagent_interleaved_tool_calls(
self,
mock_auth_default,
mock_bq_client,
mock_write_client,
mock_to_arrow_schema,
dummy_arrow_schema,
mock_asyncio_to_thread,
):
"""Two subagents call tools in same invocation — agent field is correct.
Simulates orchestrator delegating to schema_explorer first, then
image_describer, all within the same invocation.
"""
session = self._make_session()
async with managed_plugin(
PROJECT_ID, DATASET_ID, table_id=TABLE_ID
) as plugin:
await plugin._ensure_started()
mock_write_client.append_rows.reset_mock()
# Subagent 1: schema_explorer — full tool cycle
inv_ctx_1 = self._make_invocation_context(
"schema_explorer", session, invocation_id="inv-shared"
)
ctx_1 = tool_context_lib.ToolContext(invocation_context=inv_ctx_1)
tool_1 = self._make_tool("list_table_ids")
bigquery_agent_analytics_plugin.TraceManager.push_span(ctx_1, "tool")
await plugin.before_tool_callback(
tool=tool_1,
tool_args={"dataset": "analytics"},
tool_context=ctx_1,
)
await asyncio.sleep(0.01)
await plugin.after_tool_callback(
tool=tool_1,
tool_args={"dataset": "analytics"},
tool_context=ctx_1,
result={"tables": ["events", "metrics"]},
)
await asyncio.sleep(0.01)
# Subagent 2: image_describer — full tool cycle
inv_ctx_2 = self._make_invocation_context(
"image_describer", session, invocation_id="inv-shared"
)
ctx_2 = tool_context_lib.ToolContext(invocation_context=inv_ctx_2)
tool_2 = self._make_tool("describe_this_image")
bigquery_agent_analytics_plugin.TraceManager.push_span(ctx_2, "tool")
await plugin.before_tool_callback(
tool=tool_2,
tool_args={"image_uri": "https://example.com/img.jpg"},
tool_context=ctx_2,
)
await asyncio.sleep(0.01)
await plugin.after_tool_callback(
tool=tool_2,
tool_args={"image_uri": "https://example.com/img.jpg"},
tool_context=ctx_2,
result={"description": "A photo of scones"},
)
await asyncio.sleep(0.01)
rows = await _get_captured_rows_async(
mock_write_client, dummy_arrow_schema
)
assert len(rows) == 4
# schema_explorer tool events
assert rows[0]["agent"] == "schema_explorer"
assert rows[0]["event_type"] == "TOOL_STARTING"
assert rows[0]["invocation_id"] == "inv-shared"
assert json.loads(rows[0]["content"])["tool"] == "list_table_ids"
assert rows[1]["agent"] == "schema_explorer"
assert rows[1]["event_type"] == "TOOL_COMPLETED"
assert json.loads(rows[1]["content"])["result"]["tables"] == [
"events",
"metrics",
]
# image_describer tool events
assert rows[2]["agent"] == "image_describer"
assert rows[2]["event_type"] == "TOOL_STARTING"
assert rows[2]["invocation_id"] == "inv-shared"
assert json.loads(rows[2]["content"])["tool"] == "describe_this_image"
assert rows[3]["agent"] == "image_describer"
assert rows[3]["event_type"] == "TOOL_COMPLETED"
assert (
json.loads(rows[3]["content"])["result"]["description"]
== "A photo of scones"
)
# All share the same session and invocation
for row in rows:
assert row["session_id"] == "session-multi"
assert row["invocation_id"] == "inv-shared"
@pytest.mark.asyncio
async def test_multi_turn_multi_subagent_full_sequence(
self,
mock_auth_default,
mock_bq_client,
mock_write_client,
mock_to_arrow_schema,
dummy_arrow_schema,
mock_asyncio_to_thread,
):
"""Multi-turn + multi-subagent: two turns, each with different subagents.
Turn 1: user asks about data → orchestrator → schema_explorer (tool)
Turn 2: user asks about image → orchestrator → image_describer (tool)
Verifies invocation_id changes, agent name changes, session stays same.
"""
session = self._make_session()
async with managed_plugin(
PROJECT_ID, DATASET_ID, table_id=TABLE_ID
) as plugin:
await plugin._ensure_started()
mock_write_client.append_rows.reset_mock()
# ===== Turn 1: schema_explorer =====
inv_ctx_t1_orch = self._make_invocation_context(
"orchestrator", session, invocation_id="inv-t1"
)
cb_ctx_t1_orch = callback_context_lib.CallbackContext(
invocation_context=inv_ctx_t1_orch
)
# Orchestrator agent_starting
await plugin.before_agent_callback(
agent=inv_ctx_t1_orch.agent,
callback_context=cb_ctx_t1_orch,
)
await asyncio.sleep(0.01)
# Orchestrator delegates to schema_explorer
inv_ctx_t1_sub = self._make_invocation_context(
"schema_explorer", session, invocation_id="inv-t1"
)
cb_ctx_t1_sub = callback_context_lib.CallbackContext(
invocation_context=inv_ctx_t1_sub
)
tool_ctx_t1 = tool_context_lib.ToolContext(
invocation_context=inv_ctx_t1_sub
)
await plugin.before_agent_callback(
agent=inv_ctx_t1_sub.agent,
callback_context=cb_ctx_t1_sub,
)
await asyncio.sleep(0.01)
# schema_explorer calls tool
tool_1 = self._make_tool("list_dataset_ids")
bigquery_agent_analytics_plugin.TraceManager.push_span(
tool_ctx_t1, "tool"
)
await plugin.before_tool_callback(
tool=tool_1,
tool_args={"project_id": "proj"},
tool_context=tool_ctx_t1,
)
await asyncio.sleep(0.01)
await plugin.after_tool_callback(
tool=tool_1,
tool_args={"project_id": "proj"},
tool_context=tool_ctx_t1,
result={"datasets": ["ds1"]},
)
await asyncio.sleep(0.01)
# schema_explorer done
await plugin.after_agent_callback(
agent=inv_ctx_t1_sub.agent,
callback_context=cb_ctx_t1_sub,
)
await asyncio.sleep(0.01)
# Orchestrator done
await plugin.after_agent_callback(
agent=inv_ctx_t1_orch.agent,
callback_context=cb_ctx_t1_orch,
)
await asyncio.sleep(0.01)
# ===== Turn 2: image_describer =====
inv_ctx_t2_orch = self._make_invocation_context(
"orchestrator", session, invocation_id="inv-t2"
)
cb_ctx_t2_orch = callback_context_lib.CallbackContext(
invocation_context=inv_ctx_t2_orch
)
await plugin.before_agent_callback(
agent=inv_ctx_t2_orch.agent,
callback_context=cb_ctx_t2_orch,
)
await asyncio.sleep(0.01)
# Orchestrator delegates to image_describer
inv_ctx_t2_sub = self._make_invocation_context(
"image_describer", session, invocation_id="inv-t2"
)
cb_ctx_t2_sub = callback_context_lib.CallbackContext(
invocation_context=inv_ctx_t2_sub
)
tool_ctx_t2 = tool_context_lib.ToolContext(
invocation_context=inv_ctx_t2_sub
)
await plugin.before_agent_callback(
agent=inv_ctx_t2_sub.agent,
callback_context=cb_ctx_t2_sub,
)
await asyncio.sleep(0.01)
# image_describer calls tool
tool_2 = self._make_tool("describe_this_image")
bigquery_agent_analytics_plugin.TraceManager.push_span(
tool_ctx_t2, "tool"
)
await plugin.before_tool_callback(
tool=tool_2,
tool_args={"image_uri": "gs://b/img.jpg"},
tool_context=tool_ctx_t2,
)
await asyncio.sleep(0.01)
await plugin.after_tool_callback(
tool=tool_2,
tool_args={"image_uri": "gs://b/img.jpg"},
tool_context=tool_ctx_t2,
result={"desc": "Scones on a table"},
)
await asyncio.sleep(0.01)
# image_describer done
await plugin.after_agent_callback(
agent=inv_ctx_t2_sub.agent,
callback_context=cb_ctx_t2_sub,
)
await asyncio.sleep(0.01)
# Orchestrator done
await plugin.after_agent_callback(
agent=inv_ctx_t2_orch.agent,
callback_context=cb_ctx_t2_orch,
)
await asyncio.sleep(0.01)
rows = await _get_captured_rows_async(
mock_write_client, dummy_arrow_schema
)
# Turn 1: 6 rows (orch_start, sub_start, tool_start, tool_done,
# sub_done, orch_done)
# Turn 2: 6 rows (same pattern)
assert len(rows) == 12
# --- Turn 1 validation ---
t1_rows = [r for r in rows if r["invocation_id"] == "inv-t1"]
assert len(t1_rows) == 6
assert t1_rows[0]["event_type"] == "AGENT_STARTING"
assert t1_rows[0]["agent"] == "orchestrator"
assert t1_rows[1]["event_type"] == "AGENT_STARTING"
assert t1_rows[1]["agent"] == "schema_explorer"
assert t1_rows[2]["event_type"] == "TOOL_STARTING"
assert t1_rows[2]["agent"] == "schema_explorer"
assert json.loads(t1_rows[2]["content"])["tool"] == "list_dataset_ids"
assert t1_rows[3]["event_type"] == "TOOL_COMPLETED"
assert t1_rows[3]["agent"] == "schema_explorer"
assert t1_rows[4]["event_type"] == "AGENT_COMPLETED"
assert t1_rows[4]["agent"] == "schema_explorer"
assert t1_rows[5]["event_type"] == "AGENT_COMPLETED"
assert t1_rows[5]["agent"] == "orchestrator"
# --- Turn 2 validation ---
t2_rows = [r for r in rows if r["invocation_id"] == "inv-t2"]
assert len(t2_rows) == 6
assert t2_rows[0]["event_type"] == "AGENT_STARTING"
assert t2_rows[0]["agent"] == "orchestrator"
assert t2_rows[1]["event_type"] == "AGENT_STARTING"
assert t2_rows[1]["agent"] == "image_describer"
assert t2_rows[2]["event_type"] == "TOOL_STARTING"
assert t2_rows[2]["agent"] == "image_describer"
assert json.loads(t2_rows[2]["content"])["tool"] == "describe_this_image"
assert t2_rows[3]["event_type"] == "TOOL_COMPLETED"
assert t2_rows[3]["agent"] == "image_describer"
assert t2_rows[4]["event_type"] == "AGENT_COMPLETED"
assert t2_rows[4]["agent"] == "image_describer"
assert t2_rows[5]["event_type"] == "AGENT_COMPLETED"
assert t2_rows[5]["agent"] == "orchestrator"
# All rows share the same session
for row in rows:
assert row["session_id"] == "session-multi"