mirror of
https://github.com/encounter/adk-python.git
synced 2026-07-09 18:19:28 -07:00
feat: Add configuration options to BigQuery logging plugin
This change introduces BigQueryLoggerConfig to allow customization of the BigQueryAgentAnalyticsPlugin. Users can now enable/disable the plugin, specify event type allowlists and denylists, and provide a custom function to format or redact the content field before logging to BigQuery. The content logged for model and tool errors has also been enhanced. PiperOrigin-RevId: 828172241
This commit is contained in:
committed by
Copybara-Service
parent
88032cf5c5
commit
033f5a5d3f
@@ -14,12 +14,16 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
|
import dataclasses
|
||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
from datetime import timezone
|
from datetime import timezone
|
||||||
import json
|
import json
|
||||||
import logging
|
import logging
|
||||||
import threading
|
import threading
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
from typing import Callable
|
||||||
|
from typing import Dict
|
||||||
|
from typing import List
|
||||||
from typing import Optional
|
from typing import Optional
|
||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
@@ -44,6 +48,26 @@ if TYPE_CHECKING:
|
|||||||
from ..agents.invocation_context import InvocationContext
|
from ..agents.invocation_context import InvocationContext
|
||||||
|
|
||||||
|
|
||||||
|
@dataclasses.dataclass
|
||||||
|
class BigQueryLoggerConfig:
|
||||||
|
"""Configuration for the BigQueryAgentAnalyticsPlugin.
|
||||||
|
|
||||||
|
Attributes:
|
||||||
|
enabled: Whether the plugin is enabled.
|
||||||
|
event_allowlist: List of event types to log. If None, all are allowed
|
||||||
|
except those in event_denylist.
|
||||||
|
event_denylist: List of event types to not log. Takes precedence over
|
||||||
|
event_allowlist.
|
||||||
|
content_formatter: Function to format or redact the 'content' field before
|
||||||
|
logging.
|
||||||
|
"""
|
||||||
|
|
||||||
|
enabled: bool = True
|
||||||
|
event_allowlist: Optional[List[str]] = None
|
||||||
|
event_denylist: Optional[List[str]] = None
|
||||||
|
content_formatter: Optional[Callable[[Any], str]] = None
|
||||||
|
|
||||||
|
|
||||||
def _get_event_type(event: Event) -> str:
|
def _get_event_type(event: Event) -> str:
|
||||||
if event.author == "user":
|
if event.author == "user":
|
||||||
return "USER_INPUT"
|
return "USER_INPUT"
|
||||||
@@ -109,6 +133,8 @@ class BigQueryAgentAnalyticsPlugin(BasePlugin):
|
|||||||
|
|
||||||
Each log entry includes a timestamp, event type, agent name, session ID,
|
Each log entry includes a timestamp, event type, agent name, session ID,
|
||||||
invocation ID, user ID, content payload, and any error messages.
|
invocation ID, user ID, content payload, and any error messages.
|
||||||
|
|
||||||
|
Logging behavior can be customized using the BigQueryLoggerConfig.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
@@ -116,22 +142,35 @@ class BigQueryAgentAnalyticsPlugin(BasePlugin):
|
|||||||
project_id: str,
|
project_id: str,
|
||||||
dataset_id: str,
|
dataset_id: str,
|
||||||
table_id: str = "agent_events",
|
table_id: str = "agent_events",
|
||||||
|
config: Optional[BigQueryLoggerConfig] = None,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
super().__init__(name=kwargs.get("name", "BigQueryAgentAnalyticsPlugin"))
|
super().__init__(name=kwargs.get("name", "BigQueryAgentAnalyticsPlugin"))
|
||||||
self._project_id = project_id
|
self._project_id = project_id
|
||||||
self._dataset_id = dataset_id
|
self._dataset_id = dataset_id
|
||||||
self._table_id = table_id
|
self._table_id = table_id
|
||||||
|
self._config = config if config else BigQueryLoggerConfig()
|
||||||
self._bq_client: bigquery.Client | None = None
|
self._bq_client: bigquery.Client | None = None
|
||||||
self._client_init_lock = threading.Lock()
|
self._client_init_lock = threading.Lock()
|
||||||
self._init_done = False
|
self._init_done = False
|
||||||
self._init_succeeded = False
|
self._init_succeeded = False
|
||||||
|
|
||||||
|
if not self._config.enabled:
|
||||||
|
logging.info(
|
||||||
|
"BigQueryAgentAnalyticsPlugin %s is disabled by configuration.",
|
||||||
|
self.name,
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
logging.debug(
|
logging.debug(
|
||||||
"DEBUG: BigQueryAgentAnalyticsPlugin INSTANTIATED (Name: %s)", self.name
|
"DEBUG: BigQueryAgentAnalyticsPlugin INSTANTIATED (Name: %s)", self.name
|
||||||
)
|
)
|
||||||
|
|
||||||
def _ensure_initialized_sync(self):
|
def _ensure_initialized_sync(self):
|
||||||
"""Synchronous initialization of BQ client and table."""
|
"""Synchronous initialization of BQ client and table."""
|
||||||
|
if not self._config.enabled:
|
||||||
|
return
|
||||||
|
|
||||||
with self._client_init_lock:
|
with self._client_init_lock:
|
||||||
if self._init_done:
|
if self._init_done:
|
||||||
return
|
return
|
||||||
@@ -180,6 +219,39 @@ class BigQueryAgentAnalyticsPlugin(BasePlugin):
|
|||||||
self._init_succeeded = False
|
self._init_succeeded = False
|
||||||
|
|
||||||
async def _log_to_bigquery_async(self, event_dict: dict[str, Any]):
|
async def _log_to_bigquery_async(self, event_dict: dict[str, Any]):
|
||||||
|
if not self._config.enabled:
|
||||||
|
return
|
||||||
|
|
||||||
|
event_type = event_dict.get("event_type")
|
||||||
|
|
||||||
|
# Check denylist
|
||||||
|
if (
|
||||||
|
self._config.event_denylist
|
||||||
|
and event_type in self._config.event_denylist
|
||||||
|
):
|
||||||
|
return
|
||||||
|
|
||||||
|
# Check allowlist
|
||||||
|
if (
|
||||||
|
self._config.event_allowlist
|
||||||
|
and event_type not in self._config.event_allowlist
|
||||||
|
):
|
||||||
|
return
|
||||||
|
|
||||||
|
# Apply custom content formatter
|
||||||
|
if self._config.content_formatter and "content" in event_dict:
|
||||||
|
try:
|
||||||
|
event_dict["content"] = self._config.content_formatter(
|
||||||
|
event_dict["content"]
|
||||||
|
)
|
||||||
|
except Exception as e:
|
||||||
|
logging.warning(
|
||||||
|
"Error applying custom content formatter for event type %s: %s",
|
||||||
|
event_type,
|
||||||
|
e,
|
||||||
|
)
|
||||||
|
# Optionally log a generic message or the error
|
||||||
|
|
||||||
def _sync_log():
|
def _sync_log():
|
||||||
self._ensure_initialized_sync()
|
self._ensure_initialized_sync()
|
||||||
if not self._init_succeeded or not self._bq_client:
|
if not self._init_succeeded or not self._bq_client:
|
||||||
@@ -246,6 +318,7 @@ class BigQueryAgentAnalyticsPlugin(BasePlugin):
|
|||||||
"session_id": invocation_context.session.id,
|
"session_id": invocation_context.session.id,
|
||||||
"invocation_id": invocation_context.invocation_id,
|
"invocation_id": invocation_context.invocation_id,
|
||||||
"user_id": invocation_context.session.user_id,
|
"user_id": invocation_context.session.user_id,
|
||||||
|
"content": None,
|
||||||
}
|
}
|
||||||
await self._log_to_bigquery_async(event_dict)
|
await self._log_to_bigquery_async(event_dict)
|
||||||
return None
|
return None
|
||||||
@@ -286,6 +359,7 @@ class BigQueryAgentAnalyticsPlugin(BasePlugin):
|
|||||||
"session_id": invocation_context.session.id,
|
"session_id": invocation_context.session.id,
|
||||||
"invocation_id": invocation_context.invocation_id,
|
"invocation_id": invocation_context.invocation_id,
|
||||||
"user_id": invocation_context.session.user_id,
|
"user_id": invocation_context.session.user_id,
|
||||||
|
"content": None,
|
||||||
}
|
}
|
||||||
await self._log_to_bigquery_async(event_dict)
|
await self._log_to_bigquery_async(event_dict)
|
||||||
return None
|
return None
|
||||||
@@ -529,7 +603,9 @@ class BigQueryAgentAnalyticsPlugin(BasePlugin):
|
|||||||
"session_id": tool_context.session.id,
|
"session_id": tool_context.session.id,
|
||||||
"invocation_id": tool_context.invocation_id,
|
"invocation_id": tool_context.invocation_id,
|
||||||
"user_id": tool_context.session.user_id,
|
"user_id": tool_context.session.user_id,
|
||||||
"content": f"Tool Name: {tool.name}",
|
"content": (
|
||||||
|
f"Tool Name: {tool.name}, Arguments: {_format_args(tool_args)}"
|
||||||
|
),
|
||||||
"error_message": str(error),
|
"error_message": str(error),
|
||||||
}
|
}
|
||||||
await self._log_to_bigquery_async(event_dict)
|
await self._log_to_bigquery_async(event_dict)
|
||||||
|
|||||||
@@ -38,6 +38,8 @@ from google.cloud import bigquery
|
|||||||
from google.genai import types
|
from google.genai import types
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
|
BigQueryLoggerConfig = bigquery_logging_plugin.BigQueryLoggerConfig
|
||||||
|
|
||||||
|
|
||||||
class PluginTestBase:
|
class PluginTestBase:
|
||||||
"""Base class for plugin tests with common context setup."""
|
"""Base class for plugin tests with common context setup."""
|
||||||
@@ -109,14 +111,20 @@ class TestBigQueryAgentAnalyticsPlugin(PluginTestBase):
|
|||||||
)
|
)
|
||||||
self._asyncio_to_thread_patch.start()
|
self._asyncio_to_thread_patch.start()
|
||||||
|
|
||||||
self.plugin = bigquery_logging_plugin.BigQueryAgentAnalyticsPlugin(
|
self.plugin = asyncio.run(self._create_plugin())
|
||||||
|
|
||||||
|
async def _create_plugin(self, config=None):
|
||||||
|
plugin = bigquery_logging_plugin.BigQueryAgentAnalyticsPlugin(
|
||||||
project_id=self.project_id,
|
project_id=self.project_id,
|
||||||
dataset_id=self.dataset_id,
|
dataset_id=self.dataset_id,
|
||||||
table_id=self.table_id,
|
table_id=self.table_id,
|
||||||
|
config=config,
|
||||||
)
|
)
|
||||||
# Trigger lazy initialization by calling an async method once.
|
if config is None or config.enabled:
|
||||||
asyncio.run(self.plugin._log_to_bigquery_async({"event_type": "INIT"}))
|
# Trigger lazy initialization by calling an async method once.
|
||||||
self.mock_bq_client.insert_rows_json.reset_mock()
|
await plugin._log_to_bigquery_async({"event_type": "INIT"})
|
||||||
|
self.mock_bq_client.insert_rows_json.reset_mock()
|
||||||
|
return plugin
|
||||||
|
|
||||||
def _get_logged_entry(self):
|
def _get_logged_entry(self):
|
||||||
"""Helper to get the single logged entry from the mocked client."""
|
"""Helper to get the single logged entry from the mocked client."""
|
||||||
@@ -134,6 +142,98 @@ class TestBigQueryAgentAnalyticsPlugin(PluginTestBase):
|
|||||||
assert log_entry["user_id"] == "user-456"
|
assert log_entry["user_id"] == "user-456"
|
||||||
assert log_entry["timestamp"] is not None
|
assert log_entry["timestamp"] is not None
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_plugin_disabled(self):
|
||||||
|
self.mock_bq_client_cls.reset_mock()
|
||||||
|
config = BigQueryLoggerConfig(enabled=False)
|
||||||
|
plugin = await self._create_plugin(config)
|
||||||
|
user_message = types.Content(parts=[types.Part(text="Test")])
|
||||||
|
await plugin.on_user_message_callback(
|
||||||
|
invocation_context=self.invocation_context, user_message=user_message
|
||||||
|
)
|
||||||
|
self.mock_bq_client_cls.assert_not_called()
|
||||||
|
self.mock_bq_client.insert_rows_json.assert_not_called()
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_event_allowlist(self):
|
||||||
|
config = BigQueryLoggerConfig(event_allowlist=["LLM_REQUEST"])
|
||||||
|
plugin = await self._create_plugin(config)
|
||||||
|
|
||||||
|
# This should be logged
|
||||||
|
llm_request = llm_request_lib.LlmRequest(
|
||||||
|
model="gemini-pro",
|
||||||
|
contents=[types.Content(parts=[types.Part(text="Prompt")])],
|
||||||
|
)
|
||||||
|
await plugin.before_model_callback(
|
||||||
|
callback_context=self.callback_context, llm_request=llm_request
|
||||||
|
)
|
||||||
|
self.mock_bq_client.insert_rows_json.assert_called_once()
|
||||||
|
self.mock_bq_client.insert_rows_json.reset_mock()
|
||||||
|
|
||||||
|
# This should NOT be logged
|
||||||
|
user_message = types.Content(parts=[types.Part(text="What is up?")])
|
||||||
|
await plugin.on_user_message_callback(
|
||||||
|
invocation_context=self.invocation_context, user_message=user_message
|
||||||
|
)
|
||||||
|
self.mock_bq_client.insert_rows_json.assert_not_called()
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_event_denylist(self):
|
||||||
|
config = BigQueryLoggerConfig(event_denylist=["USER_MESSAGE_RECEIVED"])
|
||||||
|
plugin = await self._create_plugin(config)
|
||||||
|
|
||||||
|
# This should NOT be logged
|
||||||
|
user_message = types.Content(parts=[types.Part(text="What is up?")])
|
||||||
|
await plugin.on_user_message_callback(
|
||||||
|
invocation_context=self.invocation_context, user_message=user_message
|
||||||
|
)
|
||||||
|
self.mock_bq_client.insert_rows_json.assert_not_called()
|
||||||
|
|
||||||
|
# This should be logged
|
||||||
|
await plugin.before_run_callback(invocation_context=self.invocation_context)
|
||||||
|
self.mock_bq_client.insert_rows_json.assert_called_once()
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_content_formatter(self):
|
||||||
|
def redact_content(content):
|
||||||
|
return "[REDACTED]"
|
||||||
|
|
||||||
|
config = BigQueryLoggerConfig(content_formatter=redact_content)
|
||||||
|
plugin = await self._create_plugin(config)
|
||||||
|
|
||||||
|
user_message = types.Content(parts=[types.Part(text="Secret message")])
|
||||||
|
await plugin.on_user_message_callback(
|
||||||
|
invocation_context=self.invocation_context, user_message=user_message
|
||||||
|
)
|
||||||
|
|
||||||
|
log_entry = self._get_logged_entry()
|
||||||
|
self._assert_common_fields(log_entry, "USER_MESSAGE_RECEIVED")
|
||||||
|
assert log_entry["content"] == "[REDACTED]"
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_content_formatter_error(self):
|
||||||
|
def error_formatter(content):
|
||||||
|
raise ValueError("Formatter failed")
|
||||||
|
|
||||||
|
config = BigQueryLoggerConfig(content_formatter=error_formatter)
|
||||||
|
plugin = await self._create_plugin(config)
|
||||||
|
|
||||||
|
user_message = types.Content(parts=[types.Part(text="Test")])
|
||||||
|
with mock.patch.object(logging, "warning") as mock_log_warning:
|
||||||
|
await plugin.on_user_message_callback(
|
||||||
|
invocation_context=self.invocation_context, user_message=user_message
|
||||||
|
)
|
||||||
|
mock_log_warning.assert_called_once_with(
|
||||||
|
"Error applying custom content formatter for event type %s: %s",
|
||||||
|
"USER_MESSAGE_RECEIVED",
|
||||||
|
mock.ANY,
|
||||||
|
)
|
||||||
|
|
||||||
|
log_entry = self._get_logged_entry()
|
||||||
|
# Content should be a string, even if formatter failed
|
||||||
|
assert isinstance(log_entry["content"], str)
|
||||||
|
assert "User Content: text: 'Test'" in log_entry["content"]
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_on_user_message_callback_logs_correctly(self):
|
async def test_on_user_message_callback_logs_correctly(self):
|
||||||
user_message = types.Content(parts=[types.Part(text="What is up?")])
|
user_message = types.Content(parts=[types.Part(text="What is up?")])
|
||||||
@@ -371,8 +471,9 @@ class TestBigQueryAgentAnalyticsPlugin(PluginTestBase):
|
|||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_on_model_error_callback_logs_correctly(self):
|
async def test_on_model_error_callback_logs_correctly(self):
|
||||||
llm_request = mock.create_autospec(
|
llm_request = llm_request_lib.LlmRequest(
|
||||||
llm_request_lib.LlmRequest, instance=True
|
model="gemini-pro",
|
||||||
|
contents=[types.Content(parts=[types.Part(text="Prompt")])],
|
||||||
)
|
)
|
||||||
error = ValueError("LLM failed")
|
error = ValueError("LLM failed")
|
||||||
await self.plugin.on_model_error_callback(
|
await self.plugin.on_model_error_callback(
|
||||||
@@ -382,21 +483,26 @@ class TestBigQueryAgentAnalyticsPlugin(PluginTestBase):
|
|||||||
)
|
)
|
||||||
log_entry = self._get_logged_entry()
|
log_entry = self._get_logged_entry()
|
||||||
self._assert_common_fields(log_entry, "LLM_ERROR")
|
self._assert_common_fields(log_entry, "LLM_ERROR")
|
||||||
assert log_entry["content"] is None
|
assert (
|
||||||
|
log_entry["content"] is None
|
||||||
|
or "Request Content: " in log_entry["content"]
|
||||||
|
)
|
||||||
assert log_entry["error_message"] == "LLM failed"
|
assert log_entry["error_message"] == "LLM failed"
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_on_tool_error_callback_logs_correctly(self):
|
async def test_on_tool_error_callback_logs_correctly(self):
|
||||||
mock_tool = mock.create_autospec(base_tool_lib.BaseTool, instance=True)
|
mock_tool = mock.create_autospec(base_tool_lib.BaseTool, instance=True)
|
||||||
mock_tool.name = "MyTool"
|
mock_tool.name = "MyTool"
|
||||||
|
tool_args = {"param": "value"}
|
||||||
error = TimeoutError("Tool timed out")
|
error = TimeoutError("Tool timed out")
|
||||||
await self.plugin.on_tool_error_callback(
|
await self.plugin.on_tool_error_callback(
|
||||||
tool=mock_tool,
|
tool=mock_tool,
|
||||||
tool_args={"param": "value"},
|
tool_args=tool_args,
|
||||||
tool_context=self.tool_context,
|
tool_context=self.tool_context,
|
||||||
error=error,
|
error=error,
|
||||||
)
|
)
|
||||||
log_entry = self._get_logged_entry()
|
log_entry = self._get_logged_entry()
|
||||||
self._assert_common_fields(log_entry, "TOOL_ERROR")
|
self._assert_common_fields(log_entry, "TOOL_ERROR")
|
||||||
assert log_entry["content"] == "Tool Name: MyTool"
|
assert "Tool Name: MyTool" in log_entry["content"]
|
||||||
|
assert "Arguments: {'param': 'value'}" in log_entry["content"]
|
||||||
assert log_entry["error_message"] == "Tool timed out"
|
assert log_entry["error_message"] == "Tool timed out"
|
||||||
|
|||||||
Reference in New Issue
Block a user