diff --git a/src/agents/extensions/memory/sqlalchemy_session.py b/src/agents/extensions/memory/sqlalchemy_session.py index 81f7dcdae8..f166628c95 100644 --- a/src/agents/extensions/memory/sqlalchemy_session.py +++ b/src/agents/extensions/memory/sqlalchemy_session.py @@ -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 @@ -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 ''". +# 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].""" @@ -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 @@ -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), @@ -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, ), @@ -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 @@ -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 @@ -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() diff --git a/tests/extensions/memory/test_sqlalchemy_session.py b/tests/extensions/memory/test_sqlalchemy_session.py index b985d0a7e9..6d6446610f 100644 --- a/tests/extensions/memory/test_sqlalchemy_session.py +++ b/tests/extensions/memory/test_sqlalchemy_session.py @@ -8,7 +8,9 @@ from contextlib import asynccontextmanager from datetime import datetime, timedelta from pathlib import Path +from types import SimpleNamespace from typing import Any, cast +from unittest.mock import AsyncMock, MagicMock import pytest from openai.types.responses.response_output_message_param import ResponseOutputMessageParam @@ -17,7 +19,9 @@ ResponseReasoningItemParam, Summary, ) -from sqlalchemy import event, insert, select, text, update +from sqlalchemy import create_mock_engine, event, insert, select, text, update +from sqlalchemy.dialects import postgresql, sqlite +from sqlalchemy.exc import OperationalError, SQLAlchemyError from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine from sqlalchemy.sql import Select @@ -35,6 +39,60 @@ DB_URL = "sqlite+aiosqlite:///:memory:" +@pytest.mark.parametrize("dialect_url", ["mysql://", "mariadb://"]) +async def test_schema_create_all_compiles_for_mysql_family(dialect_url: str): + """MySQL-family schema creation includes both tables and the session-time index.""" + session = SQLAlchemySession.from_url("schema_compile", url=DB_URL) + tables = (session._sessions, session._messages) + statements: list[str] = [] + + def record(statement: Any, *args: Any, **kwargs: Any) -> None: + statements.append(str(statement.compile(dialect=engine.dialect))) + + engine = create_mock_engine(dialect_url, record) + + try: + session._metadata.create_all(engine) + for table in tables: + # CHARACTER SET must be emitted with the collation: a column given + # only a collation inherits the database character set, and the + # server rejects utf8mb4_bin against a non-utf8mb4 set with + # ERROR 1253, failing create_all() on e.g. a latin1 MySQL 5.7. + assert ( + table.c.session_id.type.compile(dialect=engine.dialect) + == "VARCHAR(190) CHARACTER SET utf8mb4 COLLATE utf8mb4_bin" + ) + finally: + await session.engine.dispose() + + assert any("CREATE TABLE agent_sessions" in statement for statement in statements) + messages_ddl = next( + statement for statement in statements if "CREATE TABLE agent_messages" in statement + ) + assert ( + "FOREIGN KEY(session_id) REFERENCES agent_sessions (session_id) ON DELETE CASCADE" + in messages_ddl + ) + assert any( + "CREATE INDEX idx_agent_messages_session_time " + "ON agent_messages (session_id, created_at)" in statement + for statement in statements + ) + + +async def test_schema_keeps_unbounded_session_ids_for_sqlite_and_postgresql(): + """SQLite and PostgreSQL retain the pre-existing unbounded string type.""" + session = SQLAlchemySession.from_url("schema_compile", url=DB_URL) + + try: + for table in (session._sessions, session._messages): + session_id_type = table.c.session_id.type + assert session_id_type.compile(dialect=postgresql.dialect()) == "VARCHAR" + assert session_id_type.compile(dialect=sqlite.dialect()) == "VARCHAR" + finally: + await session.engine.dispose() + + def _make_message_item(item_id: str, text_value: str) -> TResponseInputItem: content: ResponseOutputTextParam = { "type": "output_text", @@ -203,6 +261,331 @@ async def test_session_isolation(agent: Agent): assert "dogs" not in result.final_output.lower() +async def test_session_ids_are_case_sensitive(): + """Session IDs that differ only by case retain separate histories.""" + engine = create_async_engine(DB_URL) + upper = SQLAlchemySession("Foo", engine=engine, create_tables=True) + lower = SQLAlchemySession("foo", engine=engine, create_tables=True) + + try: + await upper.add_items([{"role": "user", "content": "upper"}]) + await lower.add_items([{"role": "user", "content": "lower"}]) + + assert await upper.get_items() == [{"role": "user", "content": "upper"}] + assert await lower.get_items() == [{"role": "user", "content": "lower"}] + finally: + await engine.dispose() + + +@pytest.mark.parametrize("dialect_name", ["mysql", "mariadb"]) +@pytest.mark.parametrize("create_tables", [True, False]) +async def test_constructor_does_not_impose_a_session_id_length_bound( + dialect_name: str, create_tables: bool +): + """The constructor leaves the actual schema authoritative for session ID length.""" + engine = MagicMock(spec=AsyncEngine) + engine.dialect = SimpleNamespace(name=dialect_name) + long_id = "a" * 191 + + session = SQLAlchemySession(long_id, engine=engine, create_tables=create_tables) + + assert session.session_id == long_id + + +class _ScalarResult: + def __init__(self, value: str | None) -> None: + self._value = value + + def scalar_one_or_none(self) -> str | None: + return self._value + + +@pytest.mark.parametrize("dialect_name", ["mysql", "mariadb"]) +async def test_validate_session_id_collation_rejects_pad_space( + dialect_name: str, +) -> None: + engine = MagicMock(spec=AsyncEngine) + engine.dialect = SimpleNamespace(name=dialect_name, is_mariadb=dialect_name == "mariadb") + session = SQLAlchemySession( + "tenant ", + engine=engine, + create_tables=True, + sessions_table="custom_sessions", + ) + conn = MagicMock() + conn.execute = AsyncMock(side_effect=[_ScalarResult("utf8mb4_bin"), _ScalarResult("PAD SPACE")]) + + with pytest.raises( + ValueError, + match=( + r"session_id 'tenant ' ends with a space, which is not distinct under the " + r"column's PAD SPACE collation 'utf8mb4_bin'; two sessions would silently " + r"share one history" + ), + ): + await session._validate_session_id_collation(conn) + + expected_calls = 1 if dialect_name == "mariadb" else 2 + assert conn.execute.await_count == expected_calls + assert conn.execute.await_args_list[0].args[1] == {"table_name": "custom_sessions"} + if dialect_name == "mysql": + assert conn.execute.await_args_list[1].args[1] == {"collation": "utf8mb4_bin"} + + +async def test_validate_session_id_collation_rejects_pad_space_on_mysql_57() -> None: + engine = MagicMock(spec=AsyncEngine) + engine.dialect = SimpleNamespace( + name="mysql", + is_mariadb=False, + _extract_error_code=lambda exc: exc.args[0].args[0], + ) + session = SQLAlchemySession("tenant ", engine=engine, create_tables=True) + conn = MagicMock() + conn.execute = AsyncMock( + side_effect=[ + _ScalarResult("utf8mb4_bin"), + OperationalError( + "SELECT PAD_ATTRIBUTE", + {}, + Exception(Exception(1054, "Unknown column PAD_ATTRIBUTE")), + ), + _ScalarResult("5.7.44"), + ] + ) + + with pytest.raises(ValueError, match="PAD SPACE collation"): + await session._validate_session_id_collation(conn) + + assert conn.execute.await_count == 3 + + +async def test_validate_session_id_collation_allows_mariadb_nopad() -> None: + engine = MagicMock(spec=AsyncEngine) + engine.dialect = SimpleNamespace(name="mysql", is_mariadb=True) + session = SQLAlchemySession("tenant ", engine=engine, create_tables=True) + conn = MagicMock() + conn.execute = AsyncMock(return_value=_ScalarResult("utf8mb4_nopad_bin")) + + await session._validate_session_id_collation(conn) + + assert conn.execute.await_count == 1 + + +async def test_validate_session_id_collation_allows_no_pad() -> None: + engine = MagicMock(spec=AsyncEngine) + engine.dialect = SimpleNamespace(name="mysql") + session = SQLAlchemySession("tenant ", engine=engine, create_tables=True) + conn = MagicMock() + conn.execute = AsyncMock( + side_effect=[_ScalarResult("utf8mb4_0900_bin"), _ScalarResult("NO PAD")] + ) + + await session._validate_session_id_collation(conn) + + +@pytest.mark.parametrize("pad_attribute", [None, "UNKNOWN"]) +async def test_validate_session_id_collation_rejects_unknown_pad_attribute( + pad_attribute: str | None, +) -> None: + engine = MagicMock(spec=AsyncEngine) + engine.dialect = SimpleNamespace(name="mysql") + session = SQLAlchemySession("tenant ", engine=engine, create_tables=True) + conn = MagicMock() + conn.execute = AsyncMock( + side_effect=[_ScalarResult("utf8mb4_0900_bin"), _ScalarResult(pad_attribute)] + ) + + with pytest.raises(RuntimeError, match="could not inspect collation padding"): + await session._validate_session_id_collation(conn) + + +@pytest.mark.parametrize("dialect_name", ["sqlite", "postgresql"]) +async def test_validate_session_id_collation_skips_non_mysql_dialects( + dialect_name: str, + monkeypatch: pytest.MonkeyPatch, +) -> None: + engine = MagicMock(spec=AsyncEngine) + engine.dialect = SimpleNamespace(name=dialect_name) + monkeypatch.setattr(SQLAlchemySession, "_configure_sqlite_engine", MagicMock()) + session = SQLAlchemySession("tenant ", engine=engine, create_tables=True) + conn = MagicMock() + conn.execute = AsyncMock() + + await session._validate_session_id_collation(conn) + + conn.execute.assert_not_awaited() + + +async def test_create_tables_false_skips_session_id_collation_validation( + monkeypatch: pytest.MonkeyPatch, +) -> None: + session = SQLAlchemySession.from_url("tenant ", url=DB_URL, create_tables=False) + validate = AsyncMock() + monkeypatch.setattr(session, "_validate_session_id_collation", validate) + + try: + await session._ensure_tables() + finally: + await session.engine.dispose() + + validate.assert_not_awaited() + + +async def test_existing_mysql_schema_validates_trailing_space_session_id( + monkeypatch: pytest.MonkeyPatch, +) -> None: + engine = MagicMock(spec=AsyncEngine) + engine.dialect = SimpleNamespace(name="mysql", is_mariadb=False) + conn = MagicMock() + connection_context = MagicMock() + connection_context.__aenter__ = AsyncMock(return_value=conn) + connection_context.__aexit__ = AsyncMock(return_value=None) + engine.connect.return_value = connection_context + session = SQLAlchemySession("tenant ", engine=engine, create_tables=False) + validate = AsyncMock() + monkeypatch.setattr(session, "_validate_session_id_collation", validate) + + await session._ensure_tables() + + validate.assert_awaited_once_with(conn) + + +async def test_existing_mysql_schema_retries_failed_pad_attribute_inspection() -> None: + engine = MagicMock(spec=AsyncEngine) + engine.dialect = SimpleNamespace(name="mysql", is_mariadb=False) + failed_conn = MagicMock() + failed_conn.execute = AsyncMock( + side_effect=[ + _ScalarResult("utf8mb4_bin"), + SQLAlchemyError("PAD_ATTRIBUTE query failed"), + _ScalarResult("8.0.36"), + ] + ) + retry_conn = MagicMock() + retry_conn.execute = AsyncMock( + side_effect=[_ScalarResult("utf8mb4_bin"), _ScalarResult("PAD SPACE")] + ) + connection_contexts = [] + for conn in (failed_conn, retry_conn): + connection_context = MagicMock() + connection_context.__aenter__ = AsyncMock(return_value=conn) + connection_context.__aexit__ = AsyncMock(return_value=None) + connection_contexts.append(connection_context) + engine.connect.side_effect = connection_contexts + session = SQLAlchemySession("tenant ", engine=engine, create_tables=False) + + with pytest.raises(SQLAlchemyError, match="PAD_ATTRIBUTE query failed"): + await session._ensure_tables() + assert session._session_id_collation_validated is False + + with pytest.raises(ValueError, match="PAD SPACE collation"): + await session._ensure_tables() + assert engine.connect.call_count == 2 + + +async def test_existing_mysql_schema_rejects_unknown_collation() -> None: + engine = MagicMock(spec=AsyncEngine) + engine.dialect = SimpleNamespace(name="mysql", is_mariadb=False) + conn = MagicMock() + conn.execute = AsyncMock(return_value=_ScalarResult(None)) + connection_context = MagicMock() + connection_context.__aenter__ = AsyncMock(return_value=conn) + connection_context.__aexit__ = AsyncMock(return_value=None) + engine.connect.return_value = connection_context + session = SQLAlchemySession("tenant ", engine=engine, create_tables=False) + + with pytest.raises(RuntimeError, match="could not inspect session_id collation"): + await session._ensure_tables() + assert session._session_id_collation_validated is False + + +async def test_existing_mysql_schema_retries_failed_collation_inspection() -> None: + engine = MagicMock(spec=AsyncEngine) + engine.dialect = SimpleNamespace(name="mysql", is_mariadb=False) + failed_conn = MagicMock() + failed_conn.execute = AsyncMock(side_effect=SQLAlchemyError("metadata query failed")) + retry_conn = MagicMock() + retry_conn.execute = AsyncMock( + side_effect=[_ScalarResult("utf8mb4_bin"), _ScalarResult("PAD SPACE")] + ) + connection_contexts = [] + for conn in (failed_conn, retry_conn): + connection_context = MagicMock() + connection_context.__aenter__ = AsyncMock(return_value=conn) + connection_context.__aexit__ = AsyncMock(return_value=None) + connection_contexts.append(connection_context) + engine.connect.side_effect = connection_contexts + session = SQLAlchemySession("tenant ", engine=engine, create_tables=False) + + with pytest.raises(SQLAlchemyError, match="metadata query failed"): + await session._ensure_tables() + assert session._session_id_collation_validated is False + + with pytest.raises(ValueError, match="PAD SPACE collation"): + await session._ensure_tables() + assert engine.connect.call_count == 2 + + +async def test_existing_schema_validation_cannot_be_skipped_during_connect() -> None: + engine = MagicMock(spec=AsyncEngine) + engine.dialect = SimpleNamespace(name="mysql", is_mariadb=False) + conn = MagicMock() + conn.execute = AsyncMock(side_effect=[_ScalarResult("utf8mb4_bin"), _ScalarResult("PAD SPACE")]) + connection_context = MagicMock() + connection_context.__aexit__ = AsyncMock(return_value=None) + engine.connect.return_value = connection_context + session = SQLAlchemySession("tenant ", engine=engine, create_tables=False) + + async def enter_with_mutated_id() -> MagicMock: + session.session_id = "tenant" + return conn + + connection_context.__aenter__ = AsyncMock(side_effect=enter_with_mutated_id) + + with pytest.raises(ValueError, match="PAD SPACE collation"): + await session._ensure_tables() + assert session._session_id_collation_validated is False + + +async def test_collation_validation_cache_does_not_hide_mutated_session_id( + monkeypatch: pytest.MonkeyPatch, +) -> None: + engine = MagicMock(spec=AsyncEngine) + engine.dialect = SimpleNamespace(name="mysql", is_mariadb=False) + conn = MagicMock() + conn.run_sync = AsyncMock() + connection_context = MagicMock() + connection_context.__aenter__ = AsyncMock(return_value=conn) + connection_context.__aexit__ = AsyncMock(return_value=None) + engine.begin.return_value = connection_context + engine.connect.return_value = connection_context + session = SQLAlchemySession("tenant", engine=engine, create_tables=True) + validate = AsyncMock() + monkeypatch.setattr(session, "_validate_session_id_collation", validate) + + await session._ensure_tables() + session.session_id = "tenant " + await session._ensure_tables() + + validate.assert_awaited_once_with(conn) + + +async def test_session_ids_keep_trailing_spaces_on_sqlite(): + """SQLite stores trailing-space IDs as distinct values.""" + engine = create_async_engine(DB_URL) + bare = SQLAlchemySession("tenant", engine=engine, create_tables=True) + padded = SQLAlchemySession("tenant ", engine=engine, create_tables=True) + + try: + await bare.add_items([{"role": "user", "content": "bare"}]) + await padded.add_items([{"role": "user", "content": "padded"}]) + + assert await bare.get_items() == [{"role": "user", "content": "bare"}] + assert await padded.get_items() == [{"role": "user", "content": "padded"}] + finally: + await engine.dispose() + + async def test_get_items_with_limit(agent: Agent): """Test the limit parameter in get_items.""" session_id = "limit_test" diff --git a/tests/extensions/memory/test_sqlalchemy_session_mysql.py b/tests/extensions/memory/test_sqlalchemy_session_mysql.py new file mode 100644 index 0000000000..e7ba8f6482 --- /dev/null +++ b/tests/extensions/memory/test_sqlalchemy_session_mysql.py @@ -0,0 +1,264 @@ +"""Opt-in MySQL / MariaDB integration tests for :class:`SQLAlchemySession`. + +The rest of ``test_sqlalchemy_session.py`` runs against SQLite and mocked +engines, which cannot establish server-side outcomes: a compiled ``CREATE +TABLE`` says nothing about whether InnoDB accepts the index, whether a +``latin1`` database default mangles 4-byte characters, or whether a +collation pads trailing spaces. Those only show up against a real server. + +These tests are therefore **opt-in** and skip unless a server is pointed at +explicitly. They need a MySQL-family async driver, which is not a project +dependency -- ``asyncmy`` is what the URLs below assume:: + + uv run --with asyncmy \\ + env OPENAI_RUN_MYSQL_SESSION_TESTS=1 \\ + OPENAI_MYSQL_SESSION_URL=mysql+asyncmy://root:pw@127.0.0.1:3306 \\ + pytest tests/extensions/memory/test_sqlalchemy_session_mysql.py + +A server can be had with:: + + docker run -d -e MYSQL_ROOT_PASSWORD=pw -p 3306:3306 mysql:8.0 + +``OPENAI_MARIADB_SESSION_URL`` points at a MariaDB server the same way; each +URL is exercised independently, so setting only one is fine. The URLs are +server-level (no trailing database name) because each test creates and drops +its own database -- that keeps runs idempotent and independent of leftover +rows from an earlier run. +""" + +from __future__ import annotations + +import os +from collections.abc import AsyncIterator +from typing import Any +from uuid import uuid4 + +import pytest + +pytest.importorskip("sqlalchemy") # Skip tests if SQLAlchemy is not installed + +from sqlalchemy import text +from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine + +from agents import TResponseInputItem +from agents.extensions.memory.sqlalchemy_session import SQLAlchemySession + +# Serial: each case creates and drops a server-side database, so parallel +# xdist workers would race on the same name space. +pytestmark = [ + pytest.mark.asyncio, + pytest.mark.serial, + pytest.mark.skipif( + os.environ.get("OPENAI_RUN_MYSQL_SESSION_TESTS") != "1", + reason="Set OPENAI_RUN_MYSQL_SESSION_TESTS=1 and OPENAI_MYSQL_SESSION_URL / " + "OPENAI_MARIADB_SESSION_URL to run MySQL-family integration tests.", + ), +] + +_SERVER_URL_ENV = { + "mysql": "OPENAI_MYSQL_SESSION_URL", + "mariadb": "OPENAI_MARIADB_SESSION_URL", +} + + +def _server_url(flavour: str) -> str: + url = os.environ.get(_SERVER_URL_ENV[flavour]) + if not url: + pytest.skip(f"Set {_SERVER_URL_ENV[flavour]} to run the {flavour} cases.") + return url.rstrip("/") + + +def _user(content: str) -> TResponseInputItem: + item: TResponseInputItem = {"role": "user", "content": content} + return item + + +def _assistant(content: str) -> TResponseInputItem: + item: TResponseInputItem = {"role": "assistant", "content": content} + return item + + +def _contents(items: list[TResponseInputItem]) -> list[str]: + return [str(item.get("content")) for item in items] + + +@pytest.fixture(params=sorted(_SERVER_URL_ENV)) +def flavour(request: pytest.FixtureRequest) -> str: + """Run every case against each configured server.""" + return str(request.param) + + +@pytest.fixture +async def engine(flavour: str) -> AsyncIterator[AsyncEngine]: + """A freshly created database, dropped again afterwards. + + Created with a deliberately non-utf8mb4 default (``latin1``) so the + session's own column-level charset is what has to carry 4-byte + characters. A server whose default is already utf8mb4 would hide that. + """ + server = _server_url(flavour) + database = f"agents_it_{uuid4().hex[:12]}" + + admin = create_async_engine(f"{server}/", isolation_level="AUTOCOMMIT") + try: + async with admin.connect() as conn: + await conn.execute( + text(f"CREATE DATABASE {database} CHARACTER SET latin1 COLLATE latin1_swedish_ci") + ) + finally: + await admin.dispose() + + db_engine = create_async_engine(f"{server}/{database}") + try: + yield db_engine + finally: + await db_engine.dispose() + admin = create_async_engine(f"{server}/", isolation_level="AUTOCOMMIT") + try: + async with admin.connect() as conn: + await conn.execute(text(f"DROP DATABASE IF EXISTS {database}")) + finally: + await admin.dispose() + + +async def _column(engine: AsyncEngine, table: str, column: str) -> Any: + async with engine.connect() as conn: + return ( + await conn.execute( + text( + "SELECT DATA_TYPE, CHARACTER_MAXIMUM_LENGTH, COLLATION_NAME " + "FROM information_schema.COLUMNS " + "WHERE TABLE_SCHEMA = DATABASE() AND TABLE_NAME = :t AND COLUMN_NAME = :c" + ), + {"t": table, "c": column}, + ) + ).one() + + +async def test_schema_is_created_on_the_server(engine: AsyncEngine) -> None: + """Tables, the session index and the message foreign key really exist. + + Read back from ``information_schema`` rather than from the emitted DDL: + an unbounded ``TEXT`` session id compiles fine but is rejected by InnoDB + as a key, which is the failure this schema avoids. + """ + session = SQLAlchemySession("schema-check", engine=engine, create_tables=True) + await session._ensure_tables() + + data_type, length, collation = await _column(engine, "agent_sessions", "session_id") + assert data_type == "varchar", "session_id must be bounded so it can be indexed" + assert length is not None and length <= 191, ( + f"session_id length {length} exceeds the utf8mb4 index-prefix limit" + ) + assert collation == "utf8mb4_bin" + + async with engine.connect() as conn: + indexes = { + row[0] + for row in ( + await conn.execute( + text( + "SELECT DISTINCT INDEX_NAME FROM information_schema.STATISTICS " + "WHERE TABLE_SCHEMA = DATABASE() AND TABLE_NAME = 'agent_messages'" + ) + ) + ).all() + } + foreign_keys = { + row[0] + for row in ( + await conn.execute( + text( + "SELECT REFERENCED_TABLE_NAME FROM information_schema.KEY_COLUMN_USAGE " + "WHERE TABLE_SCHEMA = DATABASE() AND TABLE_NAME = 'agent_messages' " + "AND REFERENCED_TABLE_NAME IS NOT NULL" + ) + ) + ).all() + } + assert indexes, "agent_messages should carry at least one index" + assert "agent_sessions" in foreign_keys + + +async def test_add_get_pop_clear_round_trip(engine: AsyncEngine) -> None: + session = SQLAlchemySession("crud", engine=engine, create_tables=True) + + await session.add_items([_user("hello"), _assistant("hi there")]) + assert _contents(await session.get_items()) == ["hello", "hi there"] + + popped = await session.pop_item() + assert popped is not None + assert str(popped.get("content")) == "hi there" + assert len(await session.get_items()) == 1 + + await session.clear_session() + assert await session.get_items() == [] + + +async def test_four_byte_characters_survive_a_latin1_database_default( + engine: AsyncEngine, +) -> None: + """The column charset, not the database default, must carry the content. + + The fixture's database is ``latin1``; an emoji stored without a + column-level utf8mb4 would be mangled or rejected rather than returned + intact. + """ + async with engine.connect() as conn: + assert (await conn.execute(text("SELECT @@character_set_database"))).scalar_one() == ( + "latin1" + ) + + session = SQLAlchemySession("unicode", engine=engine, create_tables=True) + original = "你好 😀 café" + await session.add_items([_user(original)]) + + assert _contents(await session.get_items()) == [original] + + +async def test_case_differing_session_ids_do_not_share_history(engine: AsyncEngine) -> None: + """``utf8mb4_bin`` keeps ids distinct that a ``_ci`` collation would merge.""" + lower = SQLAlchemySession("tenant", engine=engine, create_tables=True) + upper = SQLAlchemySession("Tenant", engine=engine, create_tables=True) + + await lower.add_items([_user("from-lowercase")]) + await upper.add_items([_user("from-uppercase")]) + + assert _contents(await lower.get_items()) == ["from-lowercase"] + assert _contents(await upper.get_items()) == ["from-uppercase"] + + +async def test_trailing_space_session_id_cannot_silently_share_history( + engine: AsyncEngine, +) -> None: + """A PAD SPACE collation makes ``'tenant '`` and ``'tenant'`` compare equal. + + Under MySQL 8 ``utf8mb4_bin`` is PAD SPACE, so a trailing-space id would + silently read and write another session's history. The session must + either reject it up front or keep the two genuinely separate; silently + merging them is the outcome this pins against. + """ + base = SQLAlchemySession("tenant", engine=engine, create_tables=True) + await base.add_items([_user("from-base")]) + + try: + trailing = SQLAlchemySession("tenant ", engine=engine, create_tables=True) + await trailing.add_items([_user("from-trailing-space")]) + except ValueError as exc: + assert "PAD SPACE" in str(exc) + assert _contents(await base.get_items()) == ["from-base"] + return + + # Accepted: only valid if the collation really keeps the two apart. + assert _contents(await trailing.get_items()) == ["from-trailing-space"] + assert _contents(await base.get_items()) == ["from-base"] + + +async def test_caller_managed_schema_is_usable(engine: AsyncEngine) -> None: + """``create_tables=False`` works against tables the caller already owns.""" + await SQLAlchemySession("owned", engine=engine, create_tables=True)._ensure_tables() + + session = SQLAlchemySession("owned", engine=engine, create_tables=False) + await session.add_items([_user("against pre-existing tables")]) + + assert _contents(await session.get_items()) == ["against pre-existing tables"]