fix: Fix Spanner DatabaseSessionService support

Introduce `DynamicPickleType` to handle session actions, using sqlalchemy-spanner `SpannerPickleType` when the database dialect is Spanner.

Connects to a Spanner database to store session data persistently in tables.
# Example using Spanner database:
`session_service = DatabaseSessionService(db_url="spanner+spanner:///projects/project-id/instances/instance-id/databases/database-id")`

# Example adk web command:
`adk web --session_service_uri="spanner+spanner:///projects/project-id/instances/instance-id/databases/database-id"`

PiperOrigin-RevId: 797416610
This commit is contained in:
Google Team Member
2025-08-20 12:26:53 -07:00
committed by Copybara-Service
parent c144b5347c
commit 54ed079100
2 changed files with 30 additions and 1 deletions
+1
View File
@@ -48,6 +48,7 @@ dependencies = [
"python-dateutil>=2.9.0.post0, <3.0.0", # For Vertext AI Session Service
"python-dotenv>=1.0.0, <2.0.0", # To manage environment variables
"requests>=2.32.4, <3.0.0",
"sqlalchemy-spanner>=1.14.0", # Spanner database session service
"sqlalchemy>=2.0, <3.0.0", # SQL database ORM
"starlette>=0.46.2, <1.0.0", # For FastAPI CLI
"tenacity>=8.0.0, <9.0.0", # For Retry management
@@ -18,6 +18,7 @@ from datetime import datetime
from datetime import timezone
import json
import logging
import pickle
from typing import Any
from typing import Optional
import uuid
@@ -104,6 +105,33 @@ class PreciseTimestamp(TypeDecorator):
return self.impl
class DynamicPickleType(TypeDecorator):
"""Represents a type that can be pickled."""
impl = PickleType
def load_dialect_impl(self, dialect):
if dialect.name == "spanner+spanner":
from google.cloud.sqlalchemy_spanner.sqlalchemy_spanner import SpannerPickleType
return dialect.type_descriptor(SpannerPickleType)
return self.impl
def process_bind_param(self, value, dialect):
"""Ensures the pickled value is a bytes object before passing it to the database dialect."""
if value is not None:
if dialect.name == "spanner+spanner":
return pickle.dumps(value)
return value
def process_result_value(self, value, dialect):
"""Ensures the raw bytes from the database are unpickled back into a Python object."""
if value is not None:
if dialect.name == "spanner+spanner":
return pickle.loads(value)
return value
class Base(DeclarativeBase):
"""Base class for database tables."""
@@ -209,7 +237,7 @@ class StorageEvent(Base):
PreciseTimestamp, default=func.now()
)
content: Mapped[dict[str, Any]] = mapped_column(DynamicJSON, nullable=True)
actions: Mapped[MutableDict[str, Any]] = mapped_column(PickleType)
actions: Mapped[MutableDict[str, Any]] = mapped_column(DynamicPickleType)
long_running_tool_ids_json: Mapped[Optional[str]] = mapped_column(
Text, nullable=True