diff --git a/CHANGELOG.md b/CHANGELOG.md index 8c039b445..6ee4409f8 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -80,6 +80,10 @@ Write the date in place of the "Unreleased" in the case a new version is release (e.g. streaming appends). - `distinct` on a catalog node counted nodes across the whole catalog instead of only the node's children. It is now scoped to the requested node. +- Fix the background task that purges expired Sessions and API keys. It + crashed on its first run and was never retried, so expired entries + (including the short-lived keys minted for websocket subscriptions) + could accumulate until a Principal hit the API-key limit. ## v0.2.18 (2026-09-02) diff --git a/tests/test_authn_database.py b/tests/test_authn_database.py new file mode 100644 index 000000000..7d268159d --- /dev/null +++ b/tests/test_authn_database.py @@ -0,0 +1,77 @@ +from datetime import datetime, timedelta, timezone + +import pytest +from sqlalchemy.ext.asyncio import AsyncSession, create_async_engine +from sqlalchemy.future import select + +from tiled.authn_database import orm +from tiled.authn_database.core import create_user, initialize_database, purge_expired + + +@pytest.fixture +async def db_session(): + engine = create_async_engine("sqlite+aiosqlite:///:memory:") + await initialize_database(engine) + async with AsyncSession(engine, autoflush=False, expire_on_commit=False) as session: + yield session + await engine.dispose() + + +async def test_purge_expired_api_keys(db_session): + principal = await create_user(db_session, "test", "alice") + now = datetime.now(timezone.utc) + expired = orm.APIKey( + principal_id=principal.id, + expiration_time=now - timedelta(seconds=10), + scopes=["inherit"], + first_eight="aaaaaaaa", + hashed_secret=b"0" * 32, + ) + valid = orm.APIKey( + principal_id=principal.id, + expiration_time=now + timedelta(seconds=600), + scopes=["inherit"], + first_eight="bbbbbbbb", + hashed_secret=b"1" * 32, + ) + never_expires = orm.APIKey( + principal_id=principal.id, + expiration_time=None, + scopes=["inherit"], + first_eight="cccccccc", + hashed_secret=b"2" * 32, + ) + db_session.add_all([expired, valid, never_expires]) + await db_session.commit() + + num_purged = await purge_expired(db_session, orm.APIKey) + assert num_purged == 1 + + remaining = { + key.first_eight + for key in (await db_session.execute(select(orm.APIKey))).unique().scalars() + } + assert remaining == {"bbbbbbbb", "cccccccc"} + + +async def test_purge_expired_sessions(db_session): + principal = await create_user(db_session, "test", "alice") + now = datetime.now(timezone.utc) + expired = orm.Session( + principal_id=principal.id, + expiration_time=now - timedelta(seconds=10), + state={}, + ) + valid = orm.Session( + principal_id=principal.id, + expiration_time=now + timedelta(seconds=600), + state={}, + ) + db_session.add_all([expired, valid]) + await db_session.commit() + + num_purged = await purge_expired(db_session, orm.Session) + assert num_purged == 1 + + remaining = (await db_session.execute(select(orm.Session))).unique().scalars().all() + assert [s.id for s in remaining] == [valid.id] diff --git a/tests/test_server.py b/tests/test_server.py index 6d59e62a8..8dc8be581 100644 --- a/tests/test_server.py +++ b/tests/test_server.py @@ -1,3 +1,6 @@ +import asyncio +import logging + import httpx import numpy import pytest @@ -9,7 +12,7 @@ from tiled.adapters.mapping import MapAdapter from tiled.catalog import in_memory from tiled.client import from_uri -from tiled.config import Authentication +from tiled.config import Authentication, Database from tiled.server.app import build_app, build_app_from_config from tiled.server.logging_config import LOGGING_CONFIG @@ -143,3 +146,56 @@ def test_about_reports_api_root_path(tmpdir, root_path): response = httpx.get(url + "/api/v1/") assert response.json()["meta"]["root_path"] == f"{root_path}/api" + + +@pytest.mark.asyncio +async def test_auth_database_purge_task_logs_and_retries_after_failure( + monkeypatch, caplog +): + attempts = 0 + retried = asyncio.Event() + release_purge = asyncio.Event() + purge_task = None + + async def fail_then_block(db_session, model): + nonlocal attempts, purge_task + attempts += 1 + purge_task = asyncio.current_task() + if attempts == 1: + raise RuntimeError("database unavailable") + retried.set() + await release_purge.wait() + + async def no_wait(delay): + pass + + monkeypatch.setattr("tiled.authn_database.core.purge_expired", fail_then_block) + monkeypatch.setattr("tiled.server.app.asyncio.sleep", no_wait) + + app = build_app( + MapAdapter({}), + server_settings={"database": Database(uri="sqlite:///:memory:")}, + ) + try: + with caplog.at_level(logging.WARNING, logger="tiled.server.app"): + async with app.router.lifespan_context(app): + await asyncio.wait_for(retried.wait(), timeout=1) + assert purge_task is not None + assert purge_task in app.state.tasks + assert not purge_task.done() + finally: + if purge_task is not None: + done, pending = await asyncio.wait({purge_task}, timeout=1) + assert not pending, "Purge task did not stop during lifespan shutdown." + await asyncio.gather(*done, return_exceptions=True) + + warning = next( + record + for record in caplog.records + if record.getMessage() + == "Failed to purge expired Sessions and API keys from the database." + ) + assert warning.levelno == logging.WARNING + assert warning.exc_info is not None + assert isinstance(warning.exc_info[1], RuntimeError) + assert str(warning.exc_info[1]) == "database unavailable" diff --git a/tiled/authn_database/core.py b/tiled/authn_database/core.py index de8d0364c..72fe7509b 100644 --- a/tiled/authn_database/core.py +++ b/tiled/authn_database/core.py @@ -99,9 +99,9 @@ async def purge_expired(db: AsyncSession, cls) -> int: statement = ( select(cls) .filter(cls.expiration_time.is_not(None)) - .filter(cls.expiration_time.replace(tzinfo=timezone.utc) < now) + .filter(cls.expiration_time < now) ) - result = await db.execute(statement) + result = (await db.execute(statement)).unique() for obj in result.scalars(): num_expired += 1 await db.delete(obj) diff --git a/tiled/server/app.py b/tiled/server/app.py index 9e7c2d421..08aab3791 100644 --- a/tiled/server/app.py +++ b/tiled/server/app.py @@ -841,23 +841,29 @@ async def startup_event(): async def purge_expired_sessions_and_api_keys(): PURGE_INTERVAL = 600 # seconds while True: - async with AsyncSession( - engine, autoflush=False, expire_on_commit=False - ) as db_session: - num_expired_sessions = await purge_expired( - db_session, orm.Session - ) - if num_expired_sessions: - logger.info( - f"Purged {num_expired_sessions} expired Sessions from the database." + try: + async with AsyncSession( + engine, autoflush=False, expire_on_commit=False + ) as db_session: + num_expired_sessions = await purge_expired( + db_session, orm.Session ) - num_expired_api_keys = await purge_expired( - db_session, orm.APIKey - ) - if num_expired_api_keys: - logger.info( - f"Purged {num_expired_api_keys} expired API keys from the database." + if num_expired_sessions: + logger.info( + f"Purged {num_expired_sessions} expired Sessions from the database." + ) + num_expired_api_keys = await purge_expired( + db_session, orm.APIKey ) + if num_expired_api_keys: + logger.info( + f"Purged {num_expired_api_keys} expired API keys from the database." + ) + except Exception: + logger.warning( + "Failed to purge expired Sessions and API keys from the database.", + exc_info=True, + ) await asyncio.sleep(PURGE_INTERVAL) app.state.tasks.append(