"""Psqlpy ADK store for Google Agent Development Kit session/event storage."""
import re
from typing import TYPE_CHECKING, Any, Final, Literal, NoReturn, cast
from typing_extensions import NotRequired
from sqlspec.adapters.psqlpy._typing import PsqlpyConnectionExecuteError, PsqlpyDatabaseError
from sqlspec.config import ADKConfig
from sqlspec.extensions.adk import BaseAsyncADKStore, StoredEvent, StoredSession, normalize_session_list_options
from sqlspec.extensions.adk.memory.store import BaseAsyncADKMemoryStore
from sqlspec.utils.logging import get_logger
from sqlspec.utils.type_guards import has_query_result_metadata
if TYPE_CHECKING:
from collections.abc import Sequence
from datetime import datetime, timedelta
from sqlspec.adapters.psqlpy.config import PsqlpyConfig
from sqlspec.extensions.adk import SessionOrderBy, StoredMemory
__all__ = ("PsqlpyADKConfig", "PsqlpyADKMemoryStore", "PsqlpyADKStore")
logger = get_logger("sqlspec.adapters.psqlpy.adk.store")
POSTGRES_TABLE_NOT_FOUND_SQLSTATE: Final = "42P01"
PSQLPY_STATUS_REGEX: Final[re.Pattern[str]] = re.compile(r"^([A-Z]+)(?:\s+(\d+))?\s+(\d+)$", re.IGNORECASE)
[docs]
class PsqlpyADKConfig(ADKConfig):
"""Psqlpy-specific ADK extension settings.
Use these keys inside ``extension_config["adk"]`` with the psqlpy ADK store.
"""
enable_event_generated_columns: NotRequired[bool]
"""Create PostgreSQL generated columns and indexes for common ADK event JSON paths."""
enable_covering_indexes: NotRequired[bool]
"""Add PostgreSQL INCLUDE columns to ADK event replay indexes."""
fillfactor: NotRequired[int]
"""Table fillfactor. Defaults to 80."""
autovacuum_vacuum_scale_factor: NotRequired[float]
"""Optional event-table autovacuum vacuum scale factor."""
autovacuum_analyze_scale_factor: NotRequired[float]
"""Optional event-table autovacuum analyze scale factor."""
class PsqlpyADKStore(BaseAsyncADKStore["PsqlpyConfig"]):
"""PostgreSQL ADK store using Psqlpy driver.
Implements session and event storage for Google Agent Development Kit
using PostgreSQL via the high-performance Rust-based psqlpy driver.
Events are stored as a single JSONB blob (``event_data``) alongside
indexed scalar columns for efficient querying.
Provides:
- Session state management with JSONB storage
- Full-fidelity event storage via ``event_data`` JSONB column
- Atomic ``append_event_and_update_state`` for durable session mutations
- Microsecond-precision timestamps with TIMESTAMPTZ
- Foreign key constraints with cascade delete
- GIN indexes for JSONB queries
- HOT updates with FILLFACTOR 80
Args:
config: PsqlpyConfig with extension_config["adk"] settings.
"""
__slots__ = ()
_config: "PsqlpyConfig"
def __init__(self, config: "PsqlpyConfig") -> None:
super().__init__(config)
async def create_tables(self) -> None:
if not self.create_schema_enabled:
await self.reconcile_schema()
return
async with self._config.provide_session() as driver:
await driver.execute_script(await self._sessions_table_ddl())
await driver.execute_script(await self._events_table_ddl())
await driver.execute_script(await self._app_states_table_ddl())
await driver.execute_script(await self._user_states_table_ddl())
await driver.execute_script(await self._metadata_table_ddl())
async def create_session(
self, session_id: str, app_name: str, user_id: str, state: "dict[str, Any]", owner_id: "Any | None" = None
) -> StoredSession:
async with self._config.provide_connection() as conn:
if self._owner_id_column_name:
sql = f"""
INSERT INTO {self._session_table}
(id, app_name, user_id, {self._owner_id_column_name}, state, create_time, update_time)
VALUES ($1, $2, $3, $4, $5, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)
RETURNING id, app_name, user_id, state, create_time, update_time
"""
result = await conn.fetch(sql, [session_id, app_name, user_id, owner_id, state])
else:
sql = f"""
INSERT INTO {self._session_table} (id, app_name, user_id, state, create_time, update_time)
VALUES ($1, $2, $3, $4, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)
RETURNING id, app_name, user_id, state, create_time, update_time
"""
result = await conn.fetch(sql, [session_id, app_name, user_id, state])
rows: list[dict[str, Any]] = result.result() if result else []
if not rows:
msg = "Failed to fetch created session"
raise RuntimeError(msg)
row = rows[0]
return StoredSession(
id=row["id"],
app_name=row["app_name"],
user_id=row["user_id"],
state=row["state"],
create_time=row["create_time"],
update_time=row["update_time"],
)
async def get_session(
self, app_name: str, user_id: str, session_id: str, *, renew_for: "int | timedelta | None" = None
) -> "StoredSession | None":
if renew_for is not None and self._calculate_expires_at(renew_for) is not None:
sql = f"""
UPDATE {self._session_table}
SET update_time = CURRENT_TIMESTAMP
WHERE app_name = $1 AND user_id = $2 AND id = $3
RETURNING id, app_name, user_id, state, create_time, update_time
"""
else:
sql = f"""
SELECT id, app_name, user_id, state, create_time, update_time
FROM {self._session_table}
WHERE app_name = $1 AND user_id = $2 AND id = $3
"""
try:
async with self._config.provide_connection() as conn:
result = await conn.fetch(sql, [app_name, user_id, session_id])
rows: list[dict[str, Any]] = result.result() if result else []
if not rows:
return None
row = rows[0]
return StoredSession(
id=row["id"],
app_name=row["app_name"],
user_id=row["user_id"],
state=row["state"],
create_time=row["create_time"],
update_time=row["update_time"],
)
except Exception as e:
if _is_table_missing_error(e):
return None
raise
async def update_session_state(self, app_name: str, user_id: str, session_id: str, state: "dict[str, Any]") -> None:
sql = f"""
UPDATE {self._session_table}
SET state = $1, update_time = CURRENT_TIMESTAMP
WHERE app_name = $2 AND user_id = $3 AND id = $4
"""
async with self._config.provide_connection() as conn:
await conn.execute(sql, [state, app_name, user_id, session_id])
async 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]":
column, direction, page_limit, page_offset = normalize_session_list_options(order_by, descending, limit, offset)
if page_limit == 0:
return []
params: list[Any] = [app_name]
where_clause = "app_name = $1"
if user_id is not None:
params.append(user_id)
where_clause = f"{where_clause} AND user_id = ${len(params)}"
page_clause = ""
if page_limit is not None:
params.append(page_limit)
limit_placeholder = f"${len(params)}"
params.append(page_offset)
page_clause = f"\n LIMIT {limit_placeholder} OFFSET ${len(params)}"
sql = f"""
SELECT id, app_name, user_id, state, create_time, update_time
FROM {self._session_table}
WHERE {where_clause}
ORDER BY {column} {direction}, id {direction}{page_clause}
"""
try:
async with self._config.provide_connection() as conn:
result = await conn.fetch(sql, params)
rows: list[dict[str, Any]] = result.result() if result else []
return [
StoredSession(
id=row["id"],
app_name=row["app_name"],
user_id=row["user_id"],
state=row["state"],
create_time=row["create_time"],
update_time=row["update_time"],
)
for row in rows
]
except Exception as e:
if _is_table_missing_error(e):
return []
raise
async def delete_session(self, app_name: str, user_id: str, session_id: str) -> None:
sql = f"DELETE FROM {self._session_table} WHERE app_name = $1 AND user_id = $2 AND id = $3"
async with self._config.provide_connection() as conn:
await conn.execute(sql, [app_name, user_id, session_id])
async def append_event(self, event_record: StoredEvent) -> None:
sql = f"""
INSERT INTO {self._events_table} (
id, app_name, user_id, session_id, invocation_id, timestamp, event_data
) VALUES ($1, $2, $3, $4, $5, $6, $7)
"""
async with self._config.provide_connection() as conn:
await conn.execute(
sql,
[
event_record["id"],
event_record["app_name"],
event_record["user_id"],
event_record["session_id"],
event_record["invocation_id"],
event_record["timestamp"],
event_record["event_data"],
],
)
async 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:
insert_sql = f"""
INSERT INTO {self._events_table} (
id, app_name, user_id, session_id, invocation_id, timestamp, event_data
) VALUES ($1, $2, $3, $4, $5, $6, $7)
"""
update_sql = f"""
UPDATE {self._session_table}
SET state = $1, update_time = CURRENT_TIMESTAMP
WHERE app_name = $2 AND user_id = $3 AND id = $4
RETURNING id, app_name, user_id, state, create_time, update_time
"""
app_upsert_sql = f"""
INSERT INTO {self._app_state_table} (app_name, state, update_time)
VALUES ($1, $2, CURRENT_TIMESTAMP)
ON CONFLICT (app_name) DO UPDATE SET
state = EXCLUDED.state,
update_time = CURRENT_TIMESTAMP
"""
user_upsert_sql = f"""
INSERT INTO {self._user_state_table} (app_name, user_id, state, update_time)
VALUES ($1, $2, $3, CURRENT_TIMESTAMP)
ON CONFLICT (app_name, user_id) DO UPDATE SET
state = EXCLUDED.state,
update_time = CURRENT_TIMESTAMP
"""
async with self._config.provide_connection() as conn:
try:
await conn.execute("BEGIN")
await conn.execute(
insert_sql,
[
event_record["id"],
event_record["app_name"],
event_record["user_id"],
event_record["session_id"],
event_record["invocation_id"],
event_record["timestamp"],
event_record["event_data"],
],
)
result = await conn.fetch(update_sql, [state, app_name, user_id, session_id])
rows: list[dict[str, Any]] = result.result() if result else []
if not rows:
_raise_missing_session(session_id)
if app_state is not None:
await conn.execute(app_upsert_sql, [app_name, app_state])
if user_state is not None:
await conn.execute(user_upsert_sql, [app_name, user_id, user_state])
except Exception:
await conn.execute("ROLLBACK")
raise
await conn.execute("COMMIT")
row = rows[0]
return StoredSession(
id=row["id"],
app_name=row["app_name"],
user_id=row["user_id"],
state=row["state"],
create_time=row["create_time"],
update_time=row["update_time"],
)
async def get_events(
self,
app_name: str,
user_id: str,
session_id: str,
after_timestamp: "datetime | None" = None,
limit: "int | None" = None,
) -> "list[StoredEvent]":
if limit == 0:
return []
where_clauses = ["s.app_name = $1", "s.user_id = $2", "e.session_id = $3"]
params: list[Any] = [app_name, user_id, session_id]
if after_timestamp is not None:
where_clauses.append(f"e.timestamp > ${len(params) + 1}")
params.append(after_timestamp)
where_clause = " AND ".join(where_clauses)
limit_clause = f" LIMIT ${len(params) + 1}" if limit is not None else ""
if limit is not None:
params.append(limit)
sql = f"""
SELECT e.id, e.session_id, e.invocation_id, e.timestamp, e.event_data, s.app_name, s.user_id
FROM {self._events_table} e
JOIN {self._session_table} s ON e.session_id = s.id
WHERE {where_clause}
ORDER BY e.timestamp ASC{limit_clause}
"""
try:
async with self._config.provide_connection() as conn:
result = await conn.fetch(sql, params)
rows: list[dict[str, Any]] = result.result() if result else []
return [
StoredEvent(
id=row["id"],
session_id=row["session_id"],
invocation_id=row["invocation_id"],
timestamp=row["timestamp"],
event_data=row["event_data"],
app_name=row["app_name"],
user_id=row["user_id"],
)
for row in rows
]
except Exception as e:
if _is_table_missing_error(e):
return []
raise
async def delete_expired_events(self, before: "datetime", app_name: "str | None" = None) -> int:
if app_name is not None:
count_sql = f"SELECT COUNT(*) AS count FROM {self._events_table} WHERE timestamp < $1 AND app_name = $2"
delete_sql = f"DELETE FROM {self._events_table} WHERE timestamp < $1 AND app_name = $2"
params: list[Any] = [before, app_name]
else:
count_sql = f"SELECT COUNT(*) AS count FROM {self._events_table} WHERE timestamp < $1"
delete_sql = f"DELETE FROM {self._events_table} WHERE timestamp < $1"
params = [before]
try:
async with self._config.provide_connection() as conn:
count_result = await conn.fetch(count_sql, params)
count_rows: list[dict[str, Any]] = count_result.result() if count_result else []
count = int(count_rows[0]["count"]) if count_rows else 0
await conn.execute(delete_sql, params)
return count
except Exception as e:
if _is_table_missing_error(e):
return 0
raise
async def delete_idle_sessions(self, updated_before: "datetime", app_name: "str | None" = None) -> int:
if app_name is not None:
count_sql = f"SELECT COUNT(*) AS count FROM {self._session_table} WHERE update_time < $1 AND app_name = $2"
delete_sql = f"DELETE FROM {self._session_table} WHERE update_time < $1 AND app_name = $2"
params: list[Any] = [updated_before, app_name]
else:
count_sql = f"SELECT COUNT(*) AS count FROM {self._session_table} WHERE update_time < $1"
delete_sql = f"DELETE FROM {self._session_table} WHERE update_time < $1"
params = [updated_before]
try:
async with self._config.provide_connection() as conn:
count_result = await conn.fetch(count_sql, params)
count_rows: list[dict[str, Any]] = count_result.result() if count_result else []
count = int(count_rows[0]["count"]) if count_rows else 0
await conn.execute(delete_sql, params)
return count
except Exception as e:
if _is_table_missing_error(e):
return 0
raise
async def delete_idle_user_states(self, updated_before: "datetime", app_name: "str | None" = None) -> int:
if app_name is not None:
count_sql = (
f"SELECT COUNT(*) AS count FROM {self._user_state_table} WHERE update_time < $1 AND app_name = $2"
)
delete_sql = f"DELETE FROM {self._user_state_table} WHERE update_time < $1 AND app_name = $2"
params: list[Any] = [updated_before, app_name]
else:
count_sql = f"SELECT COUNT(*) AS count FROM {self._user_state_table} WHERE update_time < $1"
delete_sql = f"DELETE FROM {self._user_state_table} WHERE update_time < $1"
params = [updated_before]
try:
async with self._config.provide_connection() as conn:
count_result = await conn.fetch(count_sql, params)
count_rows: list[dict[str, Any]] = count_result.result() if count_result else []
count = int(count_rows[0]["count"]) if count_rows else 0
await conn.execute(delete_sql, params)
return count
except Exception as e:
if _is_table_missing_error(e):
return 0
raise
async def get_app_state(self, app_name: str) -> "dict[str, Any] | None":
sql = f"SELECT state FROM {self._app_state_table} WHERE app_name = $1"
try:
async with self._config.provide_connection() as conn:
result = await conn.fetch(sql, [app_name])
rows: list[dict[str, Any]] = result.result() if result else []
return rows[0]["state"] if rows else None
except Exception as e:
if _is_table_missing_error(e):
return None
raise
async def get_user_state(self, app_name: str, user_id: str) -> "dict[str, Any] | None":
sql = f"SELECT state FROM {self._user_state_table} WHERE app_name = $1 AND user_id = $2"
try:
async with self._config.provide_connection() as conn:
result = await conn.fetch(sql, [app_name, user_id])
rows: list[dict[str, Any]] = result.result() if result else []
return rows[0]["state"] if rows else None
except Exception as e:
if _is_table_missing_error(e):
return None
raise
async def upsert_app_state(self, app_name: str, state: "dict[str, Any]") -> None:
sql = f"""
INSERT INTO {self._app_state_table} (app_name, state, update_time)
VALUES ($1, $2, CURRENT_TIMESTAMP)
ON CONFLICT (app_name) DO UPDATE SET
state = EXCLUDED.state,
update_time = CURRENT_TIMESTAMP
"""
async with self._config.provide_connection() as conn:
await conn.execute(sql, [app_name, state])
async def upsert_user_state(self, app_name: str, user_id: str, state: "dict[str, Any]") -> None:
sql = f"""
INSERT INTO {self._user_state_table} (app_name, user_id, state, update_time)
VALUES ($1, $2, $3, CURRENT_TIMESTAMP)
ON CONFLICT (app_name, user_id) DO UPDATE SET
state = EXCLUDED.state,
update_time = CURRENT_TIMESTAMP
"""
async with self._config.provide_connection() as conn:
await conn.execute(sql, [app_name, user_id, state])
async def get_metadata(self, key: str) -> "str | None":
sql = f"SELECT value FROM {self._metadata_table} WHERE key = $1"
try:
async with self._config.provide_connection() as conn:
result = await conn.fetch(sql, [key])
rows: list[dict[str, Any]] = result.result() if result else []
return rows[0]["value"] if rows else None
except Exception as e:
if _is_table_missing_error(e):
return None
raise
async def set_metadata(self, key: str, value: str) -> None:
sql = f"""
INSERT INTO {self._metadata_table} (key, value)
VALUES ($1, $2)
ON CONFLICT (key) DO UPDATE SET value = EXCLUDED.value
"""
async with self._config.provide_connection() as conn:
await conn.execute(sql, [key, value])
async def _sessions_table_ddl(self) -> str:
owner_id_line = ""
if self._owner_id_column_ddl:
owner_id_line = f",\n {self._owner_id_column_ddl}"
table_options = _postgres_table_options(_adk_config(self._config))
return f"""
CREATE TABLE IF NOT EXISTS {self._session_table} (
id VARCHAR(128) PRIMARY KEY,
app_name VARCHAR(128) NOT NULL,
user_id VARCHAR(128) NOT NULL{owner_id_line},
state JSONB NOT NULL DEFAULT '{{}}'::jsonb,
create_time TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP,
update_time TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP
){table_options};
CREATE INDEX IF NOT EXISTS idx_{self._session_table}_app_user
ON {self._session_table}(app_name, user_id);
CREATE INDEX IF NOT EXISTS idx_{self._session_table}_update_time
ON {self._session_table}(update_time DESC);
CREATE INDEX IF NOT EXISTS idx_{self._session_table}_state
ON {self._session_table} USING GIN (state)
WHERE state != '{{}}'::jsonb;
"""
async def _events_table_ddl(self) -> str:
adk_config = _adk_config(self._config)
table_options = _postgres_table_options(adk_config, include_autovacuum=True)
generated_columns, generated_indexes, covering_columns = _postgres_event_ddl_options(
adk_config, self._events_table
)
return f"""
CREATE TABLE IF NOT EXISTS {self._events_table} (
id VARCHAR(128) PRIMARY KEY,
app_name VARCHAR(128) NOT NULL,
user_id VARCHAR(128) NOT NULL,
session_id VARCHAR(128) NOT NULL,
invocation_id VARCHAR(256),
timestamp TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP,
event_data JSONB NOT NULL{generated_columns},
FOREIGN KEY (session_id) REFERENCES {self._session_table}(id) ON DELETE CASCADE
){table_options};
CREATE INDEX IF NOT EXISTS idx_{self._events_table}_session
ON {self._events_table}(session_id, timestamp ASC){covering_columns};
CREATE INDEX IF NOT EXISTS idx_{self._events_table}_app_timestamp
ON {self._events_table}(app_name, timestamp ASC);
{generated_indexes}
"""
async def _app_states_table_ddl(self) -> str:
table_options = _postgres_table_options(_adk_config(self._config))
return f"""
CREATE TABLE IF NOT EXISTS {self._app_state_table} (
app_name VARCHAR(128) PRIMARY KEY,
state JSONB NOT NULL DEFAULT '{{}}'::jsonb,
update_time TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP
){table_options};
"""
async def _user_states_table_ddl(self) -> str:
table_options = _postgres_table_options(_adk_config(self._config))
return f"""
CREATE TABLE IF NOT EXISTS {self._user_state_table} (
app_name VARCHAR(128) NOT NULL,
user_id VARCHAR(128) NOT NULL,
state JSONB NOT NULL DEFAULT '{{}}'::jsonb,
update_time TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP,
PRIMARY KEY (app_name, user_id)
){table_options};
"""
async def _metadata_table_ddl(self) -> str:
return f"""
CREATE TABLE IF NOT EXISTS {self._metadata_table} (
key VARCHAR(128) PRIMARY KEY,
value VARCHAR(512) NOT NULL
);
"""
def _drop_app_states_table_sql(self) -> str:
return f"DROP TABLE IF EXISTS {self._app_state_table}"
def _drop_user_states_table_sql(self) -> str:
return f"DROP TABLE IF EXISTS {self._user_state_table}"
def _drop_metadata_table_sql(self) -> str:
return f"DROP TABLE IF EXISTS {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 {self._events_table}",
f"DROP TABLE IF EXISTS {self._session_table}",
]
class PsqlpyADKMemoryStore(BaseAsyncADKMemoryStore["PsqlpyConfig"]):
"""PostgreSQL ADK memory store using Psqlpy driver."""
__slots__ = ()
_config: "PsqlpyConfig"
def __init__(self, config: "PsqlpyConfig") -> None:
"""Initialize Psqlpy memory store."""
super().__init__(config)
async def create_tables(self) -> None:
"""Create the memory table and indexes if they don't exist."""
if not self.create_schema_enabled:
await self.reconcile_schema()
return
if not self._enabled:
return
async with self._config.provide_session() as driver:
await driver.execute_script(await self._memory_table_ddl())
async def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: "object | None" = None) -> int:
"""Bulk insert memory entries with deduplication."""
if not self._enabled:
msg = "Memory store is disabled"
raise RuntimeError(msg)
if not entries:
return 0
inserted_count = 0
if self._owner_id_column_name:
sql = f"""
INSERT INTO {self._memory_table} (
id, session_id, app_name, user_id, scope, event_id, author,
{self._owner_id_column_name}, timestamp, content_json,
content_text, metadata_json, inserted_at
) VALUES (
$1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12, $13
)
ON CONFLICT (event_id) DO NOTHING
"""
else:
sql = f"""
INSERT INTO {self._memory_table} (
id, session_id, app_name, user_id, scope, event_id, author,
timestamp, content_json, content_text, metadata_json, inserted_at
) VALUES (
$1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12
)
ON CONFLICT (event_id) DO NOTHING
"""
async with self._config.provide_connection() as conn:
for entry in entries:
if self._owner_id_column_name:
params = [
entry["id"],
entry["session_id"],
entry["app_name"],
entry["user_id"],
entry.get("scope", "user"),
entry["event_id"],
entry["author"],
owner_id,
entry["timestamp"],
entry["content_json"],
entry["content_text"],
entry["metadata_json"],
entry["inserted_at"],
]
else:
params = [
entry["id"],
entry["session_id"],
entry["app_name"],
entry["user_id"],
entry.get("scope", "user"),
entry["event_id"],
entry["author"],
entry["timestamp"],
entry["content_json"],
entry["content_text"],
entry["metadata_json"],
entry["inserted_at"],
]
result = await conn.execute(sql, params)
rows_affected = self._extract_rows_affected(result)
if rows_affected > 0:
inserted_count += rows_affected
return inserted_count
async 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 = "Memory store is disabled"
raise RuntimeError(msg)
if not query or not query.strip():
return []
effective_limit = limit if limit is not None else self._max_results
try:
if self._use_fts:
try:
return await self._search_entries_fts(
query, app_name, user_id, effective_limit, scope_filter=scope_filter
)
except Exception as exc:
logger.warning("FTS search failed; falling back to simple search: %s", exc)
return await self._search_entries_simple(
query, app_name, user_id, effective_limit, scope_filter=scope_filter
)
except Exception as e:
if _is_psqlpy_database_error(e):
error_msg = str(e).lower()
if "does not exist" in error_msg or "relation" in error_msg:
return []
raise
async def delete_entries_by_session(self, session_id: str) -> int:
"""Delete all memory entries for a specific session."""
count_sql = f"SELECT COUNT(*) AS count FROM {self._memory_table} WHERE session_id = $1"
delete_sql = f"DELETE FROM {self._memory_table} WHERE session_id = $1"
try:
async with self._config.provide_connection() as conn:
count_result = await conn.fetch(count_sql, [session_id])
count_rows: list[dict[str, Any]] = count_result.result() if count_result else []
count = int(count_rows[0]["count"]) if count_rows else 0
await conn.execute(delete_sql, [session_id])
return count
except Exception as e:
if _is_psqlpy_database_error(e):
error_msg = str(e).lower()
if "does not exist" in error_msg or "relation" in error_msg:
return 0
raise
async def delete_entries_older_than(
self, days: int, app_name: "str | None" = None, scope: "str | None" = None
) -> int:
"""Delete memory entries older than specified days."""
clauses = ["inserted_at < (CURRENT_TIMESTAMP - ($1::int * INTERVAL '1 day'))"]
params: list[Any] = [days]
if app_name is not None:
params.append(app_name)
clauses.append(f"app_name = ${len(params)}")
if scope is not None:
params.append(scope)
clauses.append(f"scope = ${len(params)}")
where_clause = " AND ".join(clauses)
count_sql = f"""
SELECT COUNT(*) AS count FROM {self._memory_table}
WHERE {where_clause}
"""
delete_sql = f"""
DELETE FROM {self._memory_table}
WHERE {where_clause}
"""
try:
async with self._config.provide_connection() as conn:
count_result = await conn.fetch(count_sql, params)
count_rows: list[dict[str, Any]] = count_result.result() if count_result else []
count = int(count_rows[0]["count"]) if count_rows else 0
await conn.execute(delete_sql, params)
return count
except Exception as e:
if _is_psqlpy_database_error(e):
error_msg = str(e).lower()
if "does not exist" in error_msg or "relation" in error_msg:
return 0
raise
async def _memory_table_ddl(self) -> str:
"""Get PostgreSQL CREATE TABLE SQL for memory entries."""
owner_id_line = ""
if self._owner_id_column_ddl:
owner_id_line = f",\n {self._owner_id_column_ddl}"
fts_index = ""
if self._use_fts:
fts_index = f"""
CREATE INDEX IF NOT EXISTS idx_{self._memory_table}_fts
ON {self._memory_table} USING GIN (to_tsvector('english', content_text));
"""
return f"""
CREATE TABLE IF NOT EXISTS {self._memory_table} (
id VARCHAR(128) PRIMARY KEY,
session_id VARCHAR(128) NOT NULL,
app_name VARCHAR(128) NOT NULL,
user_id VARCHAR(128) NOT NULL,
scope VARCHAR(16) NOT NULL DEFAULT 'user',
event_id VARCHAR(128) NOT NULL UNIQUE,
author VARCHAR(256){owner_id_line},
timestamp TIMESTAMPTZ NOT NULL,
content_json JSONB NOT NULL,
content_text TEXT NOT NULL,
metadata_json JSONB,
inserted_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_{self._memory_table}_app_user_time
ON {self._memory_table}(app_name, user_id, timestamp DESC);
CREATE INDEX IF NOT EXISTS idx_{self._memory_table}_session
ON {self._memory_table}(session_id);
{fts_index}
"""
def _drop_memory_table_sql(self) -> "list[str]":
"""Get PostgreSQL DROP TABLE SQL statements."""
return [f"DROP TABLE IF EXISTS {self._memory_table}"]
async def _search_entries_fts(
self, query: str, app_name: str, user_id: str, limit: int, scope_filter: Literal["all", "user", "app"] = "all"
) -> "list[StoredMemory]":
if scope_filter == "all":
where_scope = "app_name = $2 AND ((scope = 'user' AND user_id = $3) OR scope = 'app')"
scope_params: list[Any] = [query, app_name, user_id]
p_lim = "$4"
elif scope_filter == "user":
where_scope = "app_name = $2 AND scope = 'user' AND user_id = $3"
scope_params = [query, app_name, user_id]
p_lim = "$4"
else:
where_scope = "app_name = $2 AND scope = 'app'"
scope_params = [query, app_name]
p_lim = "$3"
sql = f"""
SELECT id, session_id, app_name, user_id, scope, event_id, author,
timestamp, content_json, content_text, metadata_json, inserted_at,
ts_rank(to_tsvector('english', content_text), plainto_tsquery('english', $1)) as rank
FROM {self._memory_table}
WHERE {where_scope}
AND to_tsvector('english', content_text) @@ plainto_tsquery('english', $1)
ORDER BY rank DESC, timestamp DESC
LIMIT {p_lim}
"""
async with self._config.provide_connection() as conn:
result = await conn.fetch(sql, [*scope_params, limit])
rows: list[dict[str, Any]] = result.result() if result else []
return _rows_to_records(rows)
async def _search_entries_simple(
self, query: str, app_name: str, user_id: str, limit: int, scope_filter: Literal["all", "user", "app"] = "all"
) -> "list[StoredMemory]":
pattern = f"%{query}%"
if scope_filter == "all":
where_scope = "app_name = $1 AND ((scope = 'user' AND user_id = $2) OR scope = 'app')"
scope_params = [app_name, user_id, pattern]
p_lim = "$4"
p_pat = "$3"
elif scope_filter == "user":
where_scope = "app_name = $1 AND scope = 'user' AND user_id = $2"
scope_params = [app_name, user_id, pattern]
p_lim = "$4"
p_pat = "$3"
else:
where_scope = "app_name = $1 AND scope = 'app'"
scope_params = [app_name, pattern]
p_lim = "$3"
p_pat = "$2"
sql = f"""
SELECT id, session_id, app_name, user_id, scope, event_id, author,
timestamp, content_json, content_text, metadata_json, inserted_at
FROM {self._memory_table}
WHERE {where_scope}
AND content_text ILIKE {p_pat}
ORDER BY timestamp DESC
LIMIT {p_lim}
"""
async with self._config.provide_connection() as conn:
result = await conn.fetch(sql, [*scope_params, limit])
rows: list[dict[str, Any]] = result.result() if result else []
return _rows_to_records(rows)
def _extract_rows_affected(self, result: Any) -> int:
"""Extract rows affected from psqlpy result."""
try:
if has_query_result_metadata(result):
if result.tag:
return self._parse_command_tag(result.tag)
if result.status:
return self._parse_command_tag(result.status)
if isinstance(result, str):
return self._parse_command_tag(result)
except Exception as e:
logger.debug("Failed to parse psqlpy command tag: %s", e)
return -1
def _parse_command_tag(self, tag: str) -> int:
"""Parse PostgreSQL command tag to extract rows affected."""
if not tag:
return -1
match = PSQLPY_STATUS_REGEX.match(tag.strip())
if match:
command = match.group(1).upper()
if command == "INSERT" and match.group(3):
return int(match.group(3))
if command in {"UPDATE", "DELETE"} and match.group(3):
return int(match.group(3))
return -1
def _rows_to_records(rows: "list[dict[str, Any]]") -> "list[StoredMemory]":
return [
{
"id": row["id"],
"session_id": row["session_id"],
"app_name": row["app_name"],
"user_id": row["user_id"],
"scope": row.get("scope", "user"),
"event_id": row["event_id"],
"author": row["author"],
"timestamp": row["timestamp"],
"content_json": row["content_json"],
"content_text": row["content_text"],
"metadata_json": row["metadata_json"],
"inserted_at": row["inserted_at"],
"embedding": None,
}
for row in rows
]
def _is_psqlpy_database_error(exc: Exception) -> bool:
return isinstance(exc, PsqlpyDatabaseError)
def _is_table_missing_error(exc: Exception) -> bool:
if not isinstance(exc, (PsqlpyDatabaseError, PsqlpyConnectionExecuteError)):
return False
error_msg = str(exc).lower()
return "does not exist" in error_msg or "relation" in error_msg
def _adk_config(config: Any) -> PsqlpyADKConfig:
"""Return psqlpy ADK extension settings from ``extension_config["adk"]``."""
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("PsqlpyADKConfig", adk_config)
def _postgres_table_options(adk_config: PsqlpyADKConfig, *, include_autovacuum: bool = False) -> str:
options = [_postgres_fillfactor_option(adk_config)]
if include_autovacuum:
options.extend(_postgres_autovacuum_options(adk_config))
return f" WITH ({', '.join(options)})"
def _postgres_fillfactor_option(adk_config: PsqlpyADKConfig) -> str:
value = adk_config.get("fillfactor", 80)
if not isinstance(value, int) or isinstance(value, bool) or value not in range(10, 101):
msg = "extension_config['adk']['fillfactor'] must be an integer from 10 to 100"
raise ValueError(msg)
return f"fillfactor = {value}"
def _postgres_autovacuum_options(adk_config: PsqlpyADKConfig) -> "list[str]":
options: list[str] = []
for key in ("autovacuum_vacuum_scale_factor", "autovacuum_analyze_scale_factor"):
value = adk_config.get(key)
if value is not None:
if not isinstance(value, (int, float)) or isinstance(value, bool) or not 0 <= float(value) <= 1:
msg = f"extension_config['adk']['{key}'] must be a number from 0 to 1"
raise ValueError(msg)
options.append(f"{key} = {float(value):g}")
return options
def _postgres_event_ddl_options(adk_config: PsqlpyADKConfig, events_table: str) -> "tuple[str, str, str]":
generated_columns = ""
generated_indexes = ""
if adk_config.get("enable_event_generated_columns", False):
generated_columns = """,
author_gc VARCHAR(256) GENERATED ALWAYS AS (event_data->>'author') STORED,
node_path_gc TEXT GENERATED ALWAYS AS (event_data->'node_info'->>'path') STORED"""
generated_indexes = f"""
CREATE INDEX IF NOT EXISTS idx_{events_table}_author_gc
ON {events_table}(session_id, author_gc, timestamp ASC);
CREATE INDEX IF NOT EXISTS idx_{events_table}_node_path_gc
ON {events_table}(session_id, node_path_gc, timestamp ASC);
"""
covering_columns = ""
if adk_config.get("enable_covering_indexes", False):
covering_columns = " INCLUDE (invocation_id)"
return generated_columns, generated_indexes, covering_columns
def _raise_missing_session(session_id: str) -> NoReturn:
msg = f"Session {session_id} not found during append_event_and_update_state."
raise ValueError(msg)