"""mssql-python ADK stores for Google Agent Development Kit session storage."""
from datetime import datetime
from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, cast
from typing_extensions import NotRequired
from sqlspec.adapters.mssql_python._typing import MssqlPythonCursor, MssqlPythonError
from sqlspec.adapters.mssql_python.core import extract_error_number
from sqlspec.adapters.mssql_python.data_dictionary import MssqlVersionInfo
from sqlspec.config import ADKConfig
from sqlspec.extensions.adk import BaseSyncADKStore, StoredEvent, StoredSession, normalize_session_list_options
from sqlspec.extensions.adk.memory.store import BaseSyncADKMemoryStore
from sqlspec.utils.serializers import from_json, to_json
if TYPE_CHECKING:
from collections.abc import Sequence
from datetime import timedelta
from sqlspec.adapters.mssql_python.config import MssqlPythonConfig
from sqlspec.adapters.mssql_python.driver import MssqlPythonDriver
from sqlspec.extensions.adk import SessionOrderBy
from sqlspec.extensions.adk.memory._types import StoredMemory
__all__ = ("MssqlPythonADKConfig", "MssqlPythonADKMemoryStore", "MssqlPythonADKStore")
MSSQL_TABLE_NOT_FOUND_ERROR: Final[int] = 208
MSSQL_DUPLICATE_OBJECT_ERROR: Final[int] = 2714
MSSQL_DUPLICATE_INDEX_ERROR: Final[int] = 1913
MSSQL_SCHEMA: Final[str] = "dbo"
JSON_FALLBACK_COLUMN_TYPE: Final[str] = "NVARCHAR(MAX)"
JSON_NATIVE_COLUMN_TYPE: Final[str] = "JSON"
[docs]
class MssqlPythonADKConfig(ADKConfig):
"""mssql-python ADK extension settings."""
native_json: NotRequired[bool]
"""Force native SQL Server JSON columns when True, or NVARCHAR(MAX) when False."""
[docs]
class MssqlPythonADKStore(BaseSyncADKStore["MssqlPythonConfig"]):
"""Synchronous mssql-python ADK session/event store."""
connector_name: ClassVar[str] = "mssql_python"
__slots__ = ("_json_column_type", "_native_json")
[docs]
def __init__(self, config: "MssqlPythonConfig") -> None:
super().__init__(config)
adk_config = _adk_config(config)
native_json = adk_config.get("native_json")
self._native_json: bool | None = native_json if isinstance(native_json, bool) else None
self._json_column_type: str | None = None
[docs]
def create_tables(self) -> None:
"""Create ADK tables (idempotent T-SQL) and DD-gated indexes."""
if not self.create_schema_enabled:
self.reconcile_schema()
return
with self._config.provide_session() as driver:
if self._json_column_type is None:
configured = _configured_json_column_type(self._native_json)
self._json_column_type = (
configured if configured is not None else _json_column_type_from_sync_driver(driver)
)
driver.execute_script(self._sessions_table_ddl())
driver.execute_script(self._events_table_ddl())
driver.execute_script(self._app_states_table_ddl())
driver.execute_script(self._user_states_table_ddl())
driver.execute_script(self._metadata_table_ddl())
existing_indexes = _casefold_names(
driver.data_dictionary.get_indexes(driver, schema=MSSQL_SCHEMA), "index_name"
)
for index_name, index_table, columns in self._index_specs():
if _bare_name(index_name) not in existing_indexes:
driver.execute(_create_index_sql(index_table, index_name, columns))
driver.commit()
[docs]
def create_session(
self, session_id: str, app_name: str, user_id: str, state: "dict[str, Any]", owner_id: "Any | None" = None
) -> StoredSession:
"""Create a new ADK session."""
owner_column = f", {_quote_identifier(self._owner_id_column_name)}" if self._owner_id_column_name else ""
owner_param = ", ?" if self._owner_id_column_name else ""
sql = f"""
INSERT INTO {_table_ref(self._session_table)} (
id, app_name, user_id{owner_column}, state, create_time, update_time
)
OUTPUT inserted.id, inserted.app_name, inserted.user_id, inserted.state, inserted.create_time, inserted.update_time
VALUES (?, ?, ?{owner_param}, ?, SYSUTCDATETIME(), SYSUTCDATETIME())
"""
params: tuple[Any, ...]
if self._owner_id_column_name:
params = (session_id, app_name, user_id, owner_id, to_json(state))
else:
params = (session_id, app_name, user_id, to_json(state))
row = self._execute_fetchone(sql, params, commit=True)
if row is None:
msg = "Failed to fetch created session"
raise RuntimeError(msg)
return _session_record_from_row(row)
[docs]
def get_session(
self, app_name: str, user_id: str, session_id: str, *, renew_for: "int | timedelta | None" = None
) -> "StoredSession | None":
"""Return a scoped session or ``None`` if absent."""
try:
if renew_for is not None and self._calculate_expires_at(renew_for) is not None:
self._execute(
f"""
UPDATE {_table_ref(self._session_table)}
SET update_time = SYSUTCDATETIME()
WHERE app_name = ? AND user_id = ? AND id = ?
""",
(app_name, user_id, session_id),
commit=True,
)
row = self._execute_fetchone(
f"""
SELECT TOP (1) id, app_name, user_id, state, create_time, update_time
FROM {_table_ref(self._session_table)}
WHERE app_name = ? AND user_id = ? AND id = ?
""",
(app_name, user_id, session_id),
)
except MssqlPythonError as exc:
if _is_mssql_table_missing(exc):
return None
raise
return _session_record_from_row(row) if row is not None else None
[docs]
def update_session_state(self, app_name: str, user_id: str, session_id: str, state: "dict[str, Any]") -> None:
"""Replace a session's durable state."""
self._execute(
f"""
UPDATE {_table_ref(self._session_table)}
SET state = ?, update_time = SYSUTCDATETIME()
WHERE app_name = ? AND user_id = ? AND id = ?
""",
(to_json(state), app_name, user_id, session_id),
commit=True,
)
[docs]
def list_sessions(
self,
app_name: str,
user_id: "str | None" = None,
*,
order_by: "SessionOrderBy" = "update_time",
descending: bool = True,
limit: "int | None" = None,
offset: "int | None" = None,
) -> "list[StoredSession]":
"""List ADK sessions for an application, optionally scoped to a user."""
column, direction, page_limit, page_offset = normalize_session_list_options(order_by, descending, limit, offset)
if page_limit == 0:
return []
sql, params = _session_list_query(
self._session_table, app_name, user_id, column, direction, page_limit, page_offset
)
try:
rows = self._execute_fetchall(sql, params)
except MssqlPythonError as exc:
if _is_mssql_table_missing(exc):
return []
raise
return [_session_record_from_row(row) for row in rows]
[docs]
def delete_session(self, app_name: str, user_id: str, session_id: str) -> None:
"""Delete a session. Event rows cascade through the FK."""
self._execute(
f"DELETE FROM {_table_ref(self._session_table)} WHERE app_name = ? AND user_id = ? AND id = ?",
(app_name, user_id, session_id),
commit=True,
)
[docs]
def append_event(self, event_record: StoredEvent) -> None:
"""Append an event to a session."""
self._execute(_insert_event_sql(self._events_table), _event_insert_params(event_record), commit=True)
[docs]
def append_event_and_update_state(
self,
event_record: StoredEvent,
app_name: str,
user_id: str,
session_id: str,
state: "dict[str, Any]",
*,
app_state: "dict[str, Any] | None" = None,
user_state: "dict[str, Any] | None" = None,
) -> StoredSession:
"""Atomically append an event and update durable session/scoped state."""
update_sql = f"""
UPDATE {_table_ref(self._session_table)}
SET state = ?, update_time = SYSUTCDATETIME()
OUTPUT inserted.id, inserted.app_name, inserted.user_id, inserted.state, inserted.create_time, inserted.update_time
WHERE app_name = ? AND user_id = ? AND id = ?
"""
with self._config.provide_connection() as conn, MssqlPythonCursor(conn) as cursor:
try:
cursor.execute(update_sql, (to_json(state), app_name, user_id, session_id))
row = cursor.fetchone()
if row is None:
_raise_session_not_found(session_id)
cursor.execute(_insert_event_sql(self._events_table), _event_insert_params(event_record))
if app_state is not None:
cursor.execute(self._upsert_app_state_sql(), (app_name, to_json(app_state)))
if user_state is not None:
cursor.execute(self._upsert_user_state_sql(), (app_name, user_id, to_json(user_state)))
except Exception:
conn.rollback()
raise
conn.commit()
return _session_record_from_row(row)
[docs]
def get_events(
self,
app_name: str,
user_id: str,
session_id: str,
after_timestamp: "datetime | None" = None,
limit: "int | None" = None,
) -> "list[StoredEvent]":
"""Return events for a scoped session ordered by event timestamp."""
if limit == 0:
return []
sql, params = self._events_query(app_name, user_id, session_id, after_timestamp, limit)
try:
rows = self._execute_fetchall(sql, params)
except MssqlPythonError as exc:
if _is_mssql_table_missing(exc):
return []
raise
return [_event_record_from_row(row) for row in rows]
[docs]
def delete_expired_events(self, before: datetime, app_name: "str | None" = None) -> int:
"""Delete events older than ``before``."""
sql = f"DELETE FROM {_table_ref(self._events_table)} WHERE timestamp < ?"
params: list[Any] = [before]
if app_name is not None:
sql += " AND app_name = ?"
params.append(app_name)
try:
return self._execute(sql, tuple(params), commit=True)
except MssqlPythonError as exc:
if _is_mssql_table_missing(exc):
return 0
raise
[docs]
def delete_idle_sessions(self, updated_before: datetime, app_name: "str | None" = None) -> int:
"""Delete sessions whose update_time is older than ``updated_before``."""
sql = f"DELETE FROM {_table_ref(self._session_table)} WHERE update_time < ?"
params: list[Any] = [updated_before]
if app_name is not None:
sql += " AND app_name = ?"
params.append(app_name)
try:
return self._execute(sql, tuple(params), commit=True)
except MssqlPythonError as exc:
if _is_mssql_table_missing(exc):
return 0
raise
[docs]
def delete_idle_user_states(self, updated_before: datetime, app_name: "str | None" = None) -> int:
"""Delete user state rows whose update_time is older than ``updated_before``."""
sql = f"DELETE FROM {_table_ref(self._user_state_table)} WHERE update_time < ?"
params: list[Any] = [updated_before]
if app_name is not None:
sql += " AND app_name = ?"
params.append(app_name)
try:
return self._execute(sql, tuple(params), commit=True)
except MssqlPythonError as exc:
if _is_mssql_table_missing(exc):
return 0
raise
[docs]
def get_app_state(self, app_name: str) -> "dict[str, Any] | None":
"""Return app-scoped state."""
try:
row = self._execute_fetchone(
f"SELECT TOP (1) state FROM {_table_ref(self._app_state_table)} WHERE app_name = ?", (app_name,)
)
except MssqlPythonError as exc:
if _is_mssql_table_missing(exc):
return None
raise
return _json_dict(row[0]) if row is not None else None
[docs]
def get_user_state(self, app_name: str, user_id: str) -> "dict[str, Any] | None":
"""Return user-scoped state."""
try:
row = self._execute_fetchone(
f"""
SELECT TOP (1) state
FROM {_table_ref(self._user_state_table)}
WHERE app_name = ? AND user_id = ?
""",
(app_name, user_id),
)
except MssqlPythonError as exc:
if _is_mssql_table_missing(exc):
return None
raise
return _json_dict(row[0]) if row is not None else None
[docs]
def upsert_app_state(self, app_name: str, state: "dict[str, Any]") -> None:
"""Insert or replace app-scoped state."""
self._execute(self._upsert_app_state_sql(), (app_name, to_json(state)), commit=True)
[docs]
def upsert_user_state(self, app_name: str, user_id: str, state: "dict[str, Any]") -> None:
"""Insert or replace user-scoped state."""
self._execute(self._upsert_user_state_sql(), (app_name, user_id, to_json(state)), commit=True)
def _index_specs(self) -> "list[tuple[str, str, str]]":
"""Return ``(index_name, table, columns)`` specs for session and event indexes."""
return [*_sessions_index_specs(self._session_table), *_events_index_specs(self._events_table)]
def _sessions_table_ddl(self) -> str:
"""Return T-SQL DDL for the ADK session table."""
return _sessions_table_ddl(self._session_table, self._json_column_type_sync(), self._owner_id_column_ddl)
def _events_table_ddl(self) -> str:
"""Return T-SQL DDL for the ADK event table."""
return _events_table_ddl(self._events_table, self._session_table, self._json_column_type_sync())
def _app_states_table_ddl(self) -> str:
"""Return T-SQL DDL for the app-scoped state table."""
return _app_states_table_ddl(self._app_state_table, self._json_column_type_sync())
def _user_states_table_ddl(self) -> str:
"""Return T-SQL DDL for the user-scoped state table."""
return _user_states_table_ddl(self._user_state_table, self._json_column_type_sync())
def _metadata_table_ddl(self) -> str:
"""Return T-SQL DDL for the ADK metadata table."""
return _metadata_table_ddl(self._metadata_table)
def _drop_app_states_table_sql(self) -> str:
return f"DROP TABLE IF EXISTS {_table_ref(self._app_state_table)}"
def _drop_user_states_table_sql(self) -> str:
return f"DROP TABLE IF EXISTS {_table_ref(self._user_state_table)}"
def _drop_metadata_table_sql(self) -> str:
return f"DROP TABLE IF EXISTS {_table_ref(self._metadata_table)}"
def _drop_tables_sql(self) -> "list[str]":
return [
self._drop_metadata_table_sql(),
self._drop_user_states_table_sql(),
self._drop_app_states_table_sql(),
f"DROP TABLE IF EXISTS {_table_ref(self._events_table)}",
f"DROP TABLE IF EXISTS {_table_ref(self._session_table)}",
]
def _upsert_app_state_sql(self) -> str:
return _upsert_state_sql(self._app_state_table, ("app_name",), ("?",))
def _upsert_user_state_sql(self) -> str:
return _upsert_state_sql(self._user_state_table, ("app_name", "user_id"), ("?", "?"))
def _events_query(
self,
app_name: str,
user_id: str,
session_id: str,
after_timestamp: "datetime | None" = None,
limit: "int | None" = None,
) -> "tuple[str, tuple[Any, ...]]":
return _events_query(self._events_table, app_name, user_id, session_id, after_timestamp, limit)
def _json_column_type_sync(self) -> str:
if self._json_column_type is not None:
return self._json_column_type
configured = _configured_json_column_type(self._native_json)
if configured is not None:
self._json_column_type = configured
return configured
try:
with self._config.provide_session() as driver:
self._json_column_type = _json_column_type_from_sync_driver(driver)
except Exception:
return JSON_FALLBACK_COLUMN_TYPE
return self._json_column_type
def _execute_fetchone(self, sql: str, params: "tuple[Any, ...]" = (), *, commit: bool = False) -> "Any | None":
with self._config.provide_connection() as conn, MssqlPythonCursor(conn) as cursor:
cursor.execute(sql, params)
row = cursor.fetchone()
if commit:
conn.commit()
return row
def _execute_fetchall(self, sql: str, params: "tuple[Any, ...]" = ()) -> "list[Any]":
with self._config.provide_connection() as conn, MssqlPythonCursor(conn) as cursor:
cursor.execute(sql, params)
return list(cursor.fetchall())
def _execute(self, sql: str, params: "tuple[Any, ...]" = (), *, commit: bool = False) -> int:
with self._config.provide_connection() as conn, MssqlPythonCursor(conn) as cursor:
cursor.execute(sql, params)
rowcount = _cursor_rowcount(cursor)
if commit:
conn.commit()
return rowcount
[docs]
class MssqlPythonADKMemoryStore(BaseSyncADKMemoryStore["MssqlPythonConfig"]):
"""SQL Server ADK memory store using mssql-python."""
__slots__ = ()
[docs]
def __init__(self, config: "MssqlPythonConfig") -> None:
super().__init__(config)
[docs]
def create_tables(self) -> None:
"""Create the memory table (idempotent T-SQL) and DD-gated indexes."""
if not self.create_schema_enabled:
self.reconcile_schema()
return
if not self._enabled:
return
with self._config.provide_session() as driver:
driver.execute_script(self._memory_table_ddl())
existing_indexes = _casefold_names(
driver.data_dictionary.get_indexes(driver, schema=MSSQL_SCHEMA), "index_name"
)
for index_name, index_table, columns in self._memory_index_specs():
if _bare_name(index_name) not in existing_indexes:
driver.execute(_create_index_sql(index_table, index_name, columns))
driver.commit()
[docs]
def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: "object | None" = None) -> int:
"""Bulk insert memory entries with event-id deduplication."""
if not self._enabled:
msg = "ADK memory store is disabled"
raise RuntimeError(msg)
if not entries:
return 0
owner_column = f", {_quote_identifier(self._owner_id_column_name)}" if self._owner_id_column_name else ""
owner_value = ", ?" if self._owner_id_column_name else ""
sql = f"""
INSERT INTO {_table_ref(self._memory_table)} (
id, session_id, app_name, user_id, scope, event_id, author, timestamp,
content_json, content_text, metadata_json{owner_column}
)
SELECT ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?{owner_value}
WHERE NOT EXISTS (
SELECT 1 FROM {_table_ref(self._memory_table)} WITH (UPDLOCK, HOLDLOCK)
WHERE event_id = ?
);
"""
inserted = 0
with self._config.provide_connection() as conn, MssqlPythonCursor(conn) as cursor:
for entry in entries:
params: tuple[Any, ...] = (
entry["id"],
entry["session_id"],
entry["app_name"],
entry["user_id"],
entry.get("scope", "user"),
entry["event_id"],
entry.get("author"),
entry["timestamp"],
to_json(entry["content_json"]),
entry["content_text"],
to_json(entry["metadata_json"]) if entry.get("metadata_json") is not None else None,
)
if self._owner_id_column_name:
params = (*params, owner_id)
cursor.execute(sql, (*params, entry["event_id"]))
inserted += _cursor_rowcount(cursor)
conn.commit()
return inserted
[docs]
def search_entries(
self,
query: str,
app_name: str,
user_id: str,
limit: "int | None" = None,
scope_filter: Literal["all", "user", "app"] = "all",
embedding: "Sequence[float] | None" = None,
) -> "list[StoredMemory]":
"""Search memory entries by text query."""
if not self._enabled:
msg = "ADK memory store is disabled"
raise RuntimeError(msg)
limit_value = limit or self._max_results
where_scope, scope_params = _build_mssql_scope_where(app_name, user_id, scope_filter)
sql = f"""
SELECT TOP (?)
id, session_id, app_name, user_id, scope, event_id, author, timestamp,
content_json, content_text, metadata_json, inserted_at
FROM {_table_ref(self._memory_table)}
WHERE {where_scope} AND content_text LIKE ?
ORDER BY timestamp DESC
"""
rows = self._execute_fetchall(sql, (limit_value, *scope_params, f"%{query}%"))
return [_memory_record_from_row(row) for row in rows]
[docs]
def delete_entries_by_session(self, session_id: str) -> int:
"""Delete all memory entries for a specific session."""
return self._execute(
f"DELETE FROM {_table_ref(self._memory_table)} WHERE session_id = ?", (session_id,), commit=True
)
[docs]
def delete_entries_older_than(self, days: int, app_name: "str | None" = None, scope: "str | None" = None) -> int:
"""Delete memory entries older than the retention window."""
clauses = ["inserted_at < DATEADD(day, -?, SYSUTCDATETIME())"]
params: list[Any] = [days]
if app_name is not None:
clauses.append("app_name = ?")
params.append(app_name)
if scope is not None:
clauses.append("scope = ?")
params.append(scope)
where_sql = " AND ".join(clauses)
return self._execute(
f"DELETE FROM {_table_ref(self._memory_table)} WHERE {where_sql}", tuple(params), commit=True
)
def _memory_table_ddl(self) -> str:
owner_line = f",\n {self._owner_id_column_ddl}" if self._owner_id_column_ddl else ""
return f"""
IF NOT EXISTS (SELECT 1 FROM sys.tables WHERE name = N'{_escape_sql_literal(self._memory_table)}'
AND schema_id = SCHEMA_ID(N'dbo'))
BEGIN
CREATE TABLE {_table_ref(self._memory_table)} (
id NVARCHAR(128) NOT NULL,
session_id NVARCHAR(128) NOT NULL,
app_name NVARCHAR(128) NOT NULL,
user_id NVARCHAR(128) NOT NULL,
scope NVARCHAR(16) NOT NULL CONSTRAINT {_constraint_ref("df", self._memory_table, "scope")} DEFAULT N'user',
event_id NVARCHAR(128) NOT NULL,
author NVARCHAR(256) NULL,
timestamp DATETIME2(6) NOT NULL,
content_json NVARCHAR(MAX) NOT NULL,
content_text NVARCHAR(MAX) NOT NULL,
metadata_json NVARCHAR(MAX) NULL,
inserted_at DATETIME2(6) NOT NULL CONSTRAINT {_constraint_ref("df", self._memory_table, "inserted_at")}
DEFAULT SYSUTCDATETIME(){owner_line},
CONSTRAINT {_constraint_ref("pk", self._memory_table, "id")} PRIMARY KEY (id),
CONSTRAINT {_constraint_ref("uq", self._memory_table, "event_id")} UNIQUE (event_id)
);
END;
"""
def _memory_index_specs(self) -> "list[tuple[str, str, str]]":
"""Return ``(index_name, table, columns)`` specs for memory-table indexes."""
return [
(
f"idx_{self._memory_table}_app_scope_user_time",
self._memory_table,
"app_name, scope, user_id, timestamp DESC",
),
(f"idx_{self._memory_table}_scope", self._memory_table, "app_name, scope"),
(f"idx_{self._memory_table}_session", self._memory_table, "session_id"),
(f"idx_{self._memory_table}_timestamp", self._memory_table, "timestamp DESC"),
]
def _drop_memory_table_sql(self) -> "list[str]":
return [f"DROP TABLE IF EXISTS {_table_ref(self._memory_table)}"]
def _execute_fetchall(self, sql: str, params: "tuple[Any, ...]" = ()) -> "list[Any]":
with self._config.provide_connection() as conn, MssqlPythonCursor(conn) as cursor:
cursor.execute(sql, params)
return list(cursor.fetchall())
def _execute(self, sql: str, params: "tuple[Any, ...]" = (), *, commit: bool = False) -> int:
with self._config.provide_connection() as conn, MssqlPythonCursor(conn) as cursor:
cursor.execute(sql, params)
rowcount = _cursor_rowcount(cursor)
if commit:
conn.commit()
return rowcount
def _adk_config(config: Any) -> MssqlPythonADKConfig:
extension_config = getattr(config, "extension_config", {})
if not isinstance(extension_config, dict):
return {}
adk_config = extension_config.get("adk", {})
if not isinstance(adk_config, dict):
return {}
return cast("MssqlPythonADKConfig", adk_config)
def _configured_json_column_type(native_json: "bool | None") -> "str | None":
if native_json is None:
return None
if native_json is True:
return JSON_NATIVE_COLUMN_TYPE
return JSON_FALLBACK_COLUMN_TYPE
def _json_column_type_from_sync_driver(driver: "MssqlPythonDriver") -> str:
version_info = driver.data_dictionary.get_version(driver)
if isinstance(version_info, MssqlVersionInfo) and version_info.supports_native_json():
return JSON_NATIVE_COLUMN_TYPE
return JSON_FALLBACK_COLUMN_TYPE
def _sessions_table_ddl(table: str, json_column_type: str, owner_id_column_ddl: "str | None") -> str:
owner_line = f",\n {owner_id_column_ddl}" if owner_id_column_ddl else ""
return f"""
IF NOT EXISTS (SELECT 1 FROM sys.tables WHERE name = N'{_escape_sql_literal(table)}' AND schema_id = SCHEMA_ID(N'dbo'))
BEGIN
CREATE TABLE {_table_ref(table)} (
row_id UNIQUEIDENTIFIER NOT NULL CONSTRAINT {_constraint_ref("df", table, "row_id")} DEFAULT NEWSEQUENTIALID(),
id NVARCHAR(128) NOT NULL,
app_name NVARCHAR(128) NOT NULL,
user_id NVARCHAR(128) NOT NULL{owner_line},
state {json_column_type} NOT NULL,
create_time DATETIME2(6) NOT NULL CONSTRAINT {_constraint_ref("df", table, "create_time")} DEFAULT SYSUTCDATETIME(),
update_time DATETIME2(6) NOT NULL CONSTRAINT {_constraint_ref("df", table, "update_time")} DEFAULT SYSUTCDATETIME(),
CONSTRAINT {_constraint_ref("pk", table, "row_id")} PRIMARY KEY (row_id),
CONSTRAINT {_constraint_ref("uq", table, "id")} UNIQUE (id)
);
END;
"""
def _sessions_index_specs(table: str) -> "list[tuple[str, str, str]]":
return [
(f"idx_{table}_app_user", table, "app_name, user_id"),
(f"idx_{table}_update_time", table, "update_time DESC"),
]
def _events_table_ddl(table: str, session_table: str, json_column_type: str) -> str:
return f"""
IF NOT EXISTS (SELECT 1 FROM sys.tables WHERE name = N'{_escape_sql_literal(table)}' AND schema_id = SCHEMA_ID(N'dbo'))
BEGIN
CREATE TABLE {_table_ref(table)} (
row_id UNIQUEIDENTIFIER NOT NULL CONSTRAINT {_constraint_ref("df", table, "row_id")} DEFAULT NEWSEQUENTIALID(),
id NVARCHAR(128) NOT NULL,
app_name NVARCHAR(128) NOT NULL,
user_id NVARCHAR(128) NOT NULL,
session_id NVARCHAR(128) NOT NULL,
invocation_id NVARCHAR(256) NOT NULL,
timestamp DATETIME2(6) NOT NULL,
event_data {json_column_type} NOT NULL,
CONSTRAINT {_constraint_ref("pk", table, "row_id")} PRIMARY KEY (row_id),
CONSTRAINT {_constraint_ref("uq", table, "id")} UNIQUE (id),
CONSTRAINT {_constraint_ref("fk", table, "session")} FOREIGN KEY (session_id)
REFERENCES {_table_ref(session_table)}(id) ON DELETE CASCADE
);
END;
"""
def _events_index_specs(table: str) -> "list[tuple[str, str, str]]":
return [
(f"idx_{table}_scope", table, "app_name, user_id, session_id, timestamp ASC"),
(f"idx_{table}_session", table, "session_id, timestamp ASC"),
(f"idx_{table}_invocation", table, "invocation_id"),
(f"idx_{table}_timestamp", table, "timestamp ASC"),
(f"idx_{table}_app_timestamp", table, "app_name, timestamp ASC"),
]
def _app_states_table_ddl(table: str, json_column_type: str) -> str:
return f"""
IF NOT EXISTS (SELECT 1 FROM sys.tables WHERE name = N'{_escape_sql_literal(table)}' AND schema_id = SCHEMA_ID(N'dbo'))
BEGIN
CREATE TABLE {_table_ref(table)} (
app_name NVARCHAR(128) NOT NULL,
state {json_column_type} NOT NULL,
update_time DATETIME2(6) NOT NULL CONSTRAINT {_constraint_ref("df", table, "update_time")} DEFAULT SYSUTCDATETIME(),
CONSTRAINT {_constraint_ref("pk", table, "app_name")} PRIMARY KEY (app_name)
);
END;
"""
def _user_states_table_ddl(table: str, json_column_type: str) -> str:
return f"""
IF NOT EXISTS (SELECT 1 FROM sys.tables WHERE name = N'{_escape_sql_literal(table)}' AND schema_id = SCHEMA_ID(N'dbo'))
BEGIN
CREATE TABLE {_table_ref(table)} (
app_name NVARCHAR(128) NOT NULL,
user_id NVARCHAR(128) NOT NULL,
state {json_column_type} NOT NULL,
update_time DATETIME2(6) NOT NULL CONSTRAINT {_constraint_ref("df", table, "update_time")} DEFAULT SYSUTCDATETIME(),
CONSTRAINT {_constraint_ref("pk", table, "app_user")} PRIMARY KEY (app_name, user_id)
);
END;
"""
def _metadata_table_ddl(table: str) -> str:
return f"""
IF NOT EXISTS (SELECT 1 FROM sys.tables WHERE name = N'{_escape_sql_literal(table)}' AND schema_id = SCHEMA_ID(N'dbo'))
BEGIN
CREATE TABLE {_table_ref(table)} (
[key] NVARCHAR(128) NOT NULL,
value NVARCHAR(512) NOT NULL,
CONSTRAINT {_constraint_ref("pk", table, "key")} PRIMARY KEY ([key])
);
END;
"""
def _create_index_sql(table: str, index_name: str, columns: str) -> str:
return f"CREATE INDEX {_quote_identifier(index_name)} ON {_table_ref(table)} ({columns})"
def _casefold_names(rows: "list[Any]", key: str) -> "set[str]":
"""Collapse data-dictionary rows into a case-folded, schema-stripped name set."""
return {str(row.get(key, "")).rsplit(".", 1)[-1].casefold() for row in rows}
def _bare_name(name: str) -> str:
"""Return the case-folded, schema-stripped object name for membership checks."""
return name.rsplit(".", 1)[-1].casefold()
def _insert_event_sql(table: str) -> str:
return f"""
INSERT INTO {_table_ref(table)} (
id, app_name, user_id, session_id, invocation_id, timestamp, event_data
)
VALUES (?, ?, ?, ?, ?, ?, ?)
"""
def _upsert_state_sql(table: str, key_columns: "tuple[str, ...]", key_params: "tuple[str, ...]") -> str:
source_columns = ", ".join(
f"{param} AS {_quote_identifier(column)}" for column, param in zip(key_columns, key_params, strict=False)
)
source_columns = f"{source_columns}, ? AS state"
insert_columns = ", ".join(_quote_identifier(column) for column in (*key_columns, "state", "update_time"))
insert_values = ", ".join(f"source.{_quote_identifier(column)}" for column in (*key_columns, "state"))
match_clause = " AND ".join(
f"target.{_quote_identifier(column)} = source.{_quote_identifier(column)}" for column in key_columns
)
return f"""
MERGE INTO {_table_ref(table)} WITH (HOLDLOCK) AS target
USING (SELECT {source_columns}) AS source
ON ({match_clause})
WHEN MATCHED THEN
UPDATE SET state = source.state, update_time = SYSUTCDATETIME()
WHEN NOT MATCHED THEN
INSERT ({insert_columns})
VALUES ({insert_values}, SYSUTCDATETIME());
"""
def _upsert_metadata_sql(table: str) -> str:
return f"""
MERGE INTO {_table_ref(table)} WITH (HOLDLOCK) AS target
USING (SELECT ? AS [key], ? AS value) AS source
ON (target.[key] = source.[key])
WHEN MATCHED THEN
UPDATE SET value = source.value
WHEN NOT MATCHED THEN
INSERT ([key], value)
VALUES (source.[key], source.value);
"""
def _events_query(
table: str, app_name: str, user_id: str, session_id: str, after_timestamp: "datetime | None", limit: "int | None"
) -> "tuple[str, tuple[Any, ...]]":
top_clause = "TOP (?) " if limit is not None else ""
params: list[Any] = [limit] if limit is not None else []
params.extend([app_name, user_id, session_id])
after_clause = ""
if after_timestamp is not None:
after_clause = " AND timestamp > ?"
params.append(after_timestamp)
sql = f"""
SELECT {top_clause}id, app_name, user_id, session_id, invocation_id, timestamp, event_data
FROM {_table_ref(table)}
WHERE app_name = ? AND user_id = ? AND session_id = ?{after_clause}
ORDER BY timestamp ASC
"""
return sql, tuple(params)
def _event_insert_params(event_record: StoredEvent) -> "tuple[Any, ...]":
return (
event_record["id"],
event_record["app_name"],
event_record["user_id"],
event_record["session_id"],
event_record["invocation_id"],
event_record["timestamp"],
to_json(event_record["event_data"]),
)
def _session_record_from_row(row: Any) -> StoredSession:
return StoredSession(
id=row[0], app_name=row[1], user_id=row[2], state=_json_dict(row[3]), create_time=row[4], update_time=row[5]
)
def _event_record_from_row(row: Any) -> StoredEvent:
return StoredEvent(
id=row[0],
app_name=row[1],
user_id=row[2],
session_id=row[3],
invocation_id=row[4],
timestamp=row[5],
event_data=_json_dict(row[6]),
)
def _memory_record_from_row(row: Any) -> "StoredMemory":
return cast(
"StoredMemory",
{
"id": row[0],
"session_id": row[1],
"app_name": row[2],
"user_id": row[3],
"scope": row[4],
"event_id": row[5],
"author": row[6],
"timestamp": row[7],
"content_json": _json_dict(row[8]),
"content_text": row[9],
"metadata_json": _json_dict(row[10]) if row[10] is not None else None,
"inserted_at": row[11],
"embedding": None,
},
)
def _build_mssql_scope_where(
app_name: str, user_id: str, scope_filter: Literal["all", "user", "app"]
) -> "tuple[str, tuple[Any, ...]]":
if scope_filter == "all":
return "app_name = ? AND ((scope = 'user' AND user_id = ?) OR scope = 'app')", (app_name, user_id)
if scope_filter == "user":
return "app_name = ? AND scope = 'user' AND user_id = ?", (app_name, user_id)
return "app_name = ? AND scope = 'app'", (app_name,)
def _json_dict(value: Any) -> "dict[str, Any]":
if value is None:
return {}
if isinstance(value, dict):
return cast("dict[str, Any]", value)
if isinstance(value, bytearray):
value = bytes(value)
if isinstance(value, (bytes, str)):
return cast("dict[str, Any]", from_json(value))
return cast("dict[str, Any]", from_json(str(value)))
def _cursor_rowcount(cursor: Any) -> int:
rowcount = getattr(cursor, "rowcount", 0)
return rowcount if isinstance(rowcount, int) and rowcount > 0 else 0
def _is_mssql_table_missing(exc: BaseException) -> bool:
text = str(exc).lower()
return "invalid object name" in text or extract_error_number(exc) == MSSQL_TABLE_NOT_FOUND_ERROR
def _quote_identifier(identifier: str) -> str:
return f"[{identifier.replace(']', ']]')}]"
def _table_ref(table: str) -> str:
return f"{_quote_identifier(MSSQL_SCHEMA)}.{_quote_identifier(table)}"
def _constraint_ref(prefix: str, table: str, suffix: str) -> str:
return _quote_identifier(f"{prefix}_{table}_{suffix}")
def _escape_sql_literal(value: str) -> str:
return value.replace("'", "''")
def _raise_session_not_found(session_id: str) -> None:
msg = f"Session {session_id} not found during append_event_and_update_state."
raise ValueError(msg)
def _session_list_query(
session_table: str,
app_name: str,
user_id: "str | None",
column: str,
direction: str,
limit: "int | None",
offset: int,
) -> "tuple[str, tuple[Any, ...]]":
"""Return the bounded session-list query and its bound values."""
params: list[Any] = [app_name]
where_clause = "app_name = ?"
if user_id is not None:
params.append(user_id)
where_clause = f"{where_clause} AND user_id = ?"
page_clause = ""
if limit is not None:
params.extend((offset, limit))
page_clause = "\n OFFSET ? ROWS FETCH NEXT ? ROWS ONLY"
sql = f"""
SELECT id, app_name, user_id, state, create_time, update_time
FROM {_table_ref(session_table)}
WHERE {where_clause}
ORDER BY {column} {direction}, id {direction}{page_clause}
"""
return sql, tuple(params)