Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
106 changes: 101 additions & 5 deletions src/agents/extensions/memory/sqlalchemy_session.py
Original file line number Diff line number Diff line change
Expand Up @@ -48,8 +48,14 @@
text as sql_text,
update,
)
from sqlalchemy.exc import IntegrityError, OperationalError
from sqlalchemy.ext.asyncio import AsyncEngine, async_sessionmaker, create_async_engine
from sqlalchemy.dialects import mysql as mysql_dialect
from sqlalchemy.exc import IntegrityError, OperationalError, SQLAlchemyError
from sqlalchemy.ext.asyncio import (
AsyncConnection,
AsyncEngine,
async_sessionmaker,
create_async_engine,
)

from ...items import TResponseInputItem
from ...memory.session import SessionABC
Expand All @@ -62,6 +68,24 @@

_T = TypeVar("_T")

# MySQL-family dialects require a bounded VARCHAR for indexed string columns.
_MYSQL_SESSION_ID_MAX_LENGTH = 190
# ``CHARACTER SET`` is declared alongside the collation: a column given only a
# collation inherits the database character set, and the server rejects
# ``utf8mb4_bin`` against a non-utf8mb4 inherited set with
# "ERROR 1253 COLLATION 'utf8mb4_bin' is not valid for CHARACTER SET '<set>'".
# A MySQL 5.7 install defaulting to latin1 would otherwise fail in
# ``create_all()`` before either table exists.
_SESSION_ID_TYPE = String().with_variant(
mysql_dialect.VARCHAR(
_MYSQL_SESSION_ID_MAX_LENGTH,
charset="utf8mb4",
collation="utf8mb4_bin",
),
"mysql",
"mariadb",
)


class SQLAlchemySession(SessionABC):
"""SQLAlchemy implementation of [`Session`][agents.memory.session.Session]."""
Expand Down Expand Up @@ -163,7 +187,9 @@ def __init__(
'mysql+aiomysql://', or 'sqlite+aiosqlite://').
create_tables (bool, optional): Whether to automatically create the required
tables and indexes. Defaults to False for production use. Set to True for
development and testing when migrations aren't used.
development and testing when migrations aren't used. Automatically created
MySQL and MariaDB schemas store session IDs in VARCHAR(190) columns, and
session IDs longer than that are rejected only for those schemas.
sessions_table (str, optional): Override the default table name for sessions if needed.
messages_table (str, optional): Override the default table name for messages if needed.
session_settings (SessionSettings | None, optional): Session configuration settings
Expand All @@ -189,7 +215,7 @@ def __init__(
self._sessions = Table(
sessions_table,
self._metadata,
Column("session_id", String, primary_key=True),
Column("session_id", _SESSION_ID_TYPE, primary_key=True),
Column(
"created_at",
TIMESTAMP(timezone=False),
Expand All @@ -211,7 +237,7 @@ def __init__(
Column("id", Integer, primary_key=True, autoincrement=True),
Column(
"session_id",
String,
_SESSION_ID_TYPE,
ForeignKey(f"{sessions_table}.session_id", ondelete="CASCADE"),
nullable=False,
),
Expand All @@ -234,6 +260,7 @@ def __init__(
self._session_factory = async_sessionmaker(self._engine, expire_on_commit=False)

self._create_tables = create_tables
self._session_id_collation_validated = False

# ---------------------------------------------------------------------
# Convenience constructors
Expand Down Expand Up @@ -278,9 +305,74 @@ async def _deserialize_item(self, item: str) -> TResponseInputItem:
# ------------------------------------------------------------------
# Session protocol implementation
# ------------------------------------------------------------------
async def _validate_session_id_collation(self, conn: AsyncConnection) -> None:
"""Reject trailing-space IDs only when the actual MySQL collation pads spaces."""
if self._engine.dialect.name not in {"mysql", "mariadb"}:
return
collation_result = await conn.execute(
sql_text(
"SELECT COLLATION_NAME FROM information_schema.COLUMNS "
"WHERE TABLE_SCHEMA = DATABASE() AND TABLE_NAME = :table_name "
"AND COLUMN_NAME = 'session_id'"
),
{"table_name": self._sessions.name},
)
collation = collation_result.scalar_one_or_none()
if not collation:
raise RuntimeError("could not inspect session_id collation")

pad_attribute: str | None
if getattr(self._engine.dialect, "is_mariadb", False):
pad_attribute = "NO PAD" if "_nopad_" in collation.casefold() else "PAD SPACE"
else:
try:
pad_result = await conn.execute(
sql_text(
"SELECT PAD_ATTRIBUTE FROM information_schema.COLLATIONS "
"WHERE COLLATION_NAME = :collation"
),
{"collation": collation},
)
pad_attribute = pad_result.scalar_one_or_none()
except SQLAlchemyError as exc:
extract_error_code = getattr(self._engine.dialect, "_extract_error_code", None)
error_code = (
extract_error_code(getattr(exc, "orig", exc))
if callable(extract_error_code)
else None
)
if error_code != 1054:
raise
version_result = await conn.execute(sql_text("SELECT VERSION()"))
version = version_result.scalar_one_or_none()
pad_attribute = (
"PAD SPACE"
if version
and version.partition(".")[0].isdigit()
and int(version.partition(".")[0]) < 8
else None
)
if pad_attribute not in {"PAD SPACE", "NO PAD"}:
raise RuntimeError("could not inspect collation padding")

if pad_attribute == "PAD SPACE":
raise ValueError(
f"session_id {self.session_id!r} ends with a space, which is not distinct "
f"under the column's PAD SPACE collation {collation!r}; two sessions would "
"silently share one history"
)

async def _ensure_tables(self) -> None:
"""Ensure tables are created before any database operations."""
if not self._create_tables:
if (
not self._session_id_collation_validated
and self._engine.dialect.name in {"mysql", "mariadb"}
and self.session_id.endswith(" ")
):
async with self._engine.connect() as conn:
await self._validate_session_id_collation(conn)
self._session_id_collation_validated = True
return

assert self._init_lock is not None
Expand All @@ -294,6 +386,10 @@ async def _ensure_tables(self) -> None:

async with self._engine.begin() as conn:
await conn.run_sync(self._metadata.create_all)
needs_collation_validation = self.session_id.endswith(" ")
if needs_collation_validation:
await self._validate_session_id_collation(conn)
self._session_id_collation_validated = needs_collation_validation
self._create_tables = False # Only create once
finally:
self._init_lock.release()
Expand Down
Loading
Loading