Source code for sqlspec.extensions.adk.store

"""Base store class for ADK session backends."""

import inspect
import logging
from abc import ABC, abstractmethod
from datetime import datetime, timedelta, timezone
from typing import TYPE_CHECKING, Any, Final, Generic, TypeVar, cast

from sqlspec.extensions.adk._config_utils import _adk_session_store_config
from sqlspec.extensions.adk._table_utils import ensure_table_name, owner_id_column_name, unique_statements
from sqlspec.migrations.schema import SchemaTarget, ensure_schema_async, ensure_schema_sync
from sqlspec.observability import resolve_db_system
from sqlspec.utils.logging import get_logger, log_with_context
from sqlspec.utils.sync_tools import async_

if TYPE_CHECKING:
    from collections.abc import Callable

    from sqlspec.config import DatabaseConfigProtocol
    from sqlspec.extensions.adk._types import SessionOrderBy, StoredEvent, StoredSession

__all__ = ("BaseAsyncADKStore", "BaseSyncADKStore", "normalize_session_list_options")

ConfigT = TypeVar("ConfigT", bound="DatabaseConfigProtocol[Any, Any, Any]")

logger = get_logger("sqlspec.extensions.adk.store")

SESSION_ORDER_COLUMNS: Final = ("create_time", "update_time")

ADK_RESET_TABLE_PROFILES: Final = (
    ("adk_session", "adk_event", "adk_app_state", "adk_user_state", "adk_internal_metadata"),
    ("adk_session", "adk_event", "adk_app_state", "adk_user_state", "adk_metadata"),
    ("adk_sessions", "adk_events", "adk_app_states", "adk_user_states", "adk_internal_metadata"),
    ("adk_sessions", "adk_events", "adk_app_states", "adk_user_states", "adk_metadata"),
)


class _ADKStoreCommon(Generic[ConfigT]):
    """Shared non-async ADK store state and helpers."""

    if TYPE_CHECKING:
        _drop_tables_sql: "Callable[[], list[str]]"

    __slots__ = (
        "_app_state_table",
        "_config",
        "_events_table",
        "_metadata_table",
        "_owner_id_column_ddl",
        "_owner_id_column_name",
        "_session_table",
        "_user_state_table",
    )

    def __init__(self, config: ConfigT) -> None:
        """Initialize the ADK store.

        Args:
            config: SQLSpec database configuration.

        Notes:
            Reads configuration from config.extension_config["adk"]:
            - session_table: Sessions table name (default: "adk_session")
            - events_table: Events table name (default: "adk_event")
            - app_state_table: App-scoped state table name (default: "adk_app_state")
            - user_state_table: User-scoped state table name (default: "adk_user_state")
            - metadata_table: Internal metadata table name (default: "adk_internal_metadata")
            - owner_id_column: Optional owner FK column DDL (default: None)
        """
        self._config = config
        store_config = self._store_config_from_extension()
        self._session_table: str = str(store_config["session_table"])
        self._events_table: str = str(store_config["events_table"])
        self._app_state_table: str = str(store_config["app_state_table"])
        self._user_state_table: str = str(store_config["user_state_table"])
        self._metadata_table: str = str(store_config["metadata_table"])
        self._owner_id_column_ddl: str | None = store_config.get("owner_id_column")
        self._owner_id_column_name: str | None = (
            owner_id_column_name(self._owner_id_column_ddl) if self._owner_id_column_ddl else None
        )
        ensure_table_name(self._session_table)
        ensure_table_name(self._events_table)
        ensure_table_name(self._app_state_table)
        ensure_table_name(self._user_state_table)
        ensure_table_name(self._metadata_table)

    @property
    def config(self) -> ConfigT:
        """Return the database configuration."""
        return self._config

    @property
    def session_table(self) -> str:
        """Return the sessions table name."""
        return self._session_table

    @property
    def events_table(self) -> str:
        """Return the events table name."""
        return self._events_table

    @property
    def app_state_table(self) -> str:
        """Return the app-scoped state table name."""
        return self._app_state_table

    @property
    def user_state_table(self) -> str:
        """Return the user-scoped state table name."""
        return self._user_state_table

    @property
    def metadata_table(self) -> str:
        """Return the ADK metadata table name."""
        return self._metadata_table

    @property
    def owner_id_column_ddl(self) -> "str | None":
        """Return the full owner ID column DDL (or None if not configured)."""
        return self._owner_id_column_ddl

    @property
    def owner_id_column_name(self) -> "str | None":
        """Return the owner ID column name only (or None if not configured)."""
        return self._owner_id_column_name

    @property
    def create_schema_enabled(self) -> bool:
        """Return whether adapter-level table creation should run."""
        manage_schema, create_schema = self._schema_management_flags()
        return manage_schema and create_schema

    def _reset_drop_tables_sql(self) -> "list[str]":
        """Return all table drops needed before recreating the clean-break schema."""
        statements = list(self._drop_tables_sql())
        for table_profile in ADK_RESET_TABLE_PROFILES:
            statements.extend(self._drop_sql_for_table_profile(table_profile))
        return unique_statements(statements)

    def _store_config_from_extension(self) -> "dict[str, Any]":
        """Extract ADK store configuration from config.extension_config.

        Returns:
            Dict with ADK table names and optionally owner_id_column.
        """
        return dict(_adk_session_store_config(self._config))

    def _schema_management_flags(self) -> "tuple[bool, bool]":
        """Return automatic-management and missing-table creation flags."""
        extension_config = cast("dict[str, Any]", self._config.extension_config)
        settings = cast("dict[str, Any]", extension_config.get("adk", {}))
        return bool(settings.get("manage_schema", True)), bool(settings.get("create_schema", True))

    def _calculate_expires_at(self, expires_in: "int | timedelta | None") -> "datetime | None":
        """Calculate expiration timestamp from expires_in.

        Args:
            expires_in: Seconds or timedelta until expiration.

        Returns:
            UTC datetime of expiration, or None if no expiration.
        """
        if expires_in is None:
            return None

        expires_in_seconds = int(expires_in.total_seconds()) if isinstance(expires_in, timedelta) else expires_in

        if expires_in_seconds <= 0:
            return None

        return datetime.now(timezone.utc) + timedelta(seconds=expires_in_seconds)

    def _drop_sql_for_table_profile(self, table_profile: "tuple[str, str, str, str, str]") -> "list[str]":
        session_table, events_table, app_state_table, user_state_table, metadata_table = table_profile
        current_session_table = self._session_table
        current_events_table = self._events_table
        current_app_state_table = self._app_state_table
        current_user_state_table = self._user_state_table
        current_table = self._metadata_table
        self._session_table = session_table
        self._events_table = events_table
        self._app_state_table = app_state_table
        self._user_state_table = user_state_table
        self._metadata_table = metadata_table
        try:
            return list(self._drop_tables_sql())
        finally:
            self._session_table = current_session_table
            self._events_table = current_events_table
            self._app_state_table = current_app_state_table
            self._user_state_table = current_user_state_table
            self._metadata_table = current_table

    def _log_tables_created(self) -> None:
        log_with_context(
            logger,
            logging.DEBUG,
            "adk.tables.ready",
            db_system=resolve_db_system(type(self).__name__),
            session_table=self._session_table,
            events_table=self._events_table,
        )

    def _log_tables_dropped(self) -> None:
        log_with_context(
            logger,
            logging.DEBUG,
            "adk.tables.dropped",
            db_system=resolve_db_system(type(self).__name__),
            session_table=self._session_table,
            events_table=self._events_table,
        )

    def _log_tables_recreated(self) -> None:
        log_with_context(
            logger,
            logging.DEBUG,
            "adk.tables.recreated",
            db_system=resolve_db_system(type(self).__name__),
            session_table=self._session_table,
            events_table=self._events_table,
        )


class BaseAsyncADKStore(_ADKStoreCommon[ConfigT], ABC):
    """Base class for async SQLSpec-backed ADK session stores.

    Implements storage operations for Google ADK sessions and events using
    SQLSpec database adapters with async/await.

    This abstract base class provides common functionality for all database-specific
    store implementations including:
    - Connection management via SQLSpec configs
    - Table name validation
    - Session and event CRUD operations

    Subclasses must implement dialect-specific SQL queries and will be created
    in each adapter directory (e.g., sqlspec/adapters/asyncpg/adk/store.py).

    Args:
        config: SQLSpec database configuration with extension_config["adk"] settings.

    Notes:
        Configuration is read from config.extension_config["adk"]:
        - session_table: Sessions table name (default: "adk_session")
        - events_table: Events table name (default: "adk_event")
        - app_state_table: App-scoped state table name (default: "adk_app_state")
        - user_state_table: User-scoped state table name (default: "adk_user_state")
        - metadata_table: Internal metadata table name (default: "adk_internal_metadata")
        - owner_id_column: Optional owner FK column DDL (default: None)
    """

    __slots__ = ()

    async def create_tables(self) -> None:
        """Create the sessions and events tables if they don't exist."""
        raise NotImplementedError
async def prepare_schema_async(self, driver: Any) -> None: """Prepare adapter-specific schema decisions with an asynchronous driver."""