fix(tests): isolate database and module state

This commit is contained in:
Alexandre Teixeira
2026-10-03 03:46:52 +01:00
parent 3468ad36d7
commit 9fd6919ee9
8 changed files with 244 additions and 76 deletions
+8 -9
View File
@@ -8,15 +8,12 @@ import pytest
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
# Importing core.database below runs init_db() at import time, and its default
# (sqlite:///./data/app.db) can't be opened in a clean worktree because SQLite
# won't create the missing ./data parent dir - pytest then dies during
# collection, before any test module loads. Default to an in-memory DB for the
# test session so collection is deterministic and writes no repo-local
# artifacts. An explicit DATABASE_URL (a real test/CI database) is preserved.
# This only unblocks collection/import-time init; it does not provide a shared
# file-backed DB across processes - tests needing that must set DATABASE_URL.
os.environ.setdefault("DATABASE_URL", "sqlite:///:memory:")
# core.database initializes its engine during import. Always isolate that
# bootstrap from an inherited developer DATABASE_URL, before collection can
# import it. Tests needing files own their disposable databases explicitly.
# Restore the caller's environment when pytest's configuration is torn down.
_database_environment = pytest.MonkeyPatch()
_database_environment.setenv("DATABASE_URL", "sqlite:///:memory:")
# Pre-import real heavy modules BEFORE any test file's module-level stubs can
# replace them with MagicMock. Some test files (e.g. test_llm_core_sanitize_*)
@@ -101,6 +98,8 @@ def pytest_configure(config):
unknown-mark warnings still surface genuine typos outside the taxonomy. This
only registers marker names; it imports no production module.
"""
config.add_cleanup(_database_environment.undo)
import pathlib
from tests._taxonomy import discover_markers
+51
View File
@@ -0,0 +1,51 @@
"""Disposable databases for tests that exercise the real ORM and session manager."""
from contextlib import contextmanager
from tempfile import TemporaryDirectory
import pytest
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
from sqlalchemy.pool import NullPool
@contextmanager
def disposable_database(tmp_path):
"""Own the database file and engine; keep the canonical ORM classes intact."""
import core.database as database
with TemporaryDirectory(prefix="database-", dir=tmp_path) as directory:
engine = create_engine(
f"sqlite:///{directory}/test.db",
connect_args={"check_same_thread": False},
poolclass=NullPool,
)
try:
database.Base.metadata.create_all(engine)
yield sessionmaker(bind=engine, autoflush=False, autocommit=False)
finally:
engine.dispose()
@contextmanager
def isolated_session_database(tmp_path):
"""Temporarily bind the real manager and database aliases without reloading.
Reloading core.database changes its ORM classes while existing imports keep
the old classes and factories. Patch only resource bindings instead, and
undo them before disposing the owned engine and removing its files.
"""
import core.database as database
import core.session_manager as manager
import src.database as compatibility_database
with disposable_database(tmp_path) as factory:
engine = factory.kw["bind"]
with pytest.MonkeyPatch.context() as patcher:
patcher.setenv("DATABASE_URL", str(engine.url))
for module in (database, compatibility_database):
patcher.setattr(module, "DATABASE_URL", str(engine.url))
patcher.setattr(module, "engine", engine)
patcher.setattr(module, "SessionLocal", factory)
patcher.setattr(manager, "SessionLocal", factory)
yield manager.SessionManager(), database
+9 -9
View File
@@ -5,23 +5,23 @@ check-in for one user pulled EVERY user's calendar events (summaries,
locations) into their digest — a cross-tenant leak. Ownership lives on
CalendarCal.owner; the query must join it, like routes/calendar_routes.
"""
import tempfile
import uuid
import sys
from datetime import datetime
import pytest
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
from sqlalchemy.pool import NullPool
from tests.helpers.database import disposable_database
import core.database as cdb
from core.database import CalendarEvent, CalendarCal
from src.task_scheduler import _checkin_calendar_events
_TMPDB = tempfile.NamedTemporaryFile(suffix=".db", delete=False)
_ENGINE = create_engine(f"sqlite:///{_TMPDB.name}", connect_args={"check_same_thread": False}, poolclass=NullPool)
cdb.Base.metadata.create_all(_ENGINE)
_TS = sessionmaker(bind=_ENGINE, autoflush=False, autocommit=False)
@pytest.fixture(autouse=True)
def _digest_database(tmp_path):
with disposable_database(tmp_path) as factory:
with pytest.MonkeyPatch.context() as patcher:
patcher.setattr(sys.modules[__name__], "_TS", factory, raising=False)
yield
def _seed():
+135
View File
@@ -0,0 +1,135 @@
"""Guard database ownership at the helper and actual pytest lifecycle seams."""
import os
from pathlib import Path
import subprocess
import sys
import textwrap
import pytest
from tests.helpers.database import isolated_session_database
@pytest.mark.parametrize("fail_inside", [False, True])
def test_session_database_restores_bindings_and_removes_files(tmp_path, fail_inside):
import core.database as database
import core.session_manager as manager
import src.database as compatibility_database
from core.models import ChatMessage
modules = (database, compatibility_database, manager)
names = ("DATABASE_URL", "engine", "SessionLocal", "Base", "Session", "ChatMessage")
before = [{name: getattr(module, name) for name in names if hasattr(module, name)}
for module in modules]
previous_url = os.environ.get("DATABASE_URL")
saved_manager_class = manager.SessionManager
saved_db_session = manager.DbSession
saved_db_message = manager.DbChatMessage
listener = database.set_sqlite_pragma
class IntentionalFailure(Exception):
pass
try:
with isolated_session_database(tmp_path) as (sm, db_module):
owned_path = Path(db_module.engine.url.database)
assert owned_path.is_file()
assert owned_path.is_relative_to(tmp_path)
assert db_module.Session is saved_db_session
assert db_module.ChatMessage is saved_db_message
assert sm.__class__ is saved_manager_class
assert compatibility_database.SessionLocal is manager.SessionLocal
sm.create_session(session_id="owned", name="t", endpoint_url="x",
model="m", rag=False, owner="alice")
sm.add_message("owned", ChatMessage("user", "keep"))
sm.add_message("owned", ChatMessage("user", "remove"))
assert sm.truncate_messages("owned", 1)
with db_module.SessionLocal() as db:
assert db.query(saved_db_message).filter_by(session_id="owned").count() == 1
assert db.query(saved_db_session).filter_by(id="owned").one().message_count == 1
if fail_inside:
raise IntentionalFailure
except IntentionalFailure:
assert fail_inside
assert os.environ.get("DATABASE_URL") == previous_url
for module, bindings in zip(modules, before):
for name, value in bindings.items():
assert getattr(module, name) is value
assert database.set_sqlite_pragma is listener
assert manager.DbSession is saved_db_session
assert manager.DbChatMessage is saved_db_message
assert manager.SessionManager is saved_manager_class
assert not owned_path.parent.exists()
with isolated_session_database(tmp_path) as (sm, db_module):
assert db_module.engine.url.database != str(owned_path)
with pytest.raises(KeyError, match="Session owned not found"):
sm.get_session("owned")
with db_module.SessionLocal() as db:
assert db.query(saved_db_session).count() == 0
@pytest.mark.parametrize("truncation_first", [True, False])
def test_actual_tests_restore_process_state_and_ignore_inherited_database(tmp_path, truncation_first):
# An inherited developer URL must never be opened, even during collection.
inherited_db = tmp_path / "developer.db"
sentinel = b"a developer database must not be opened or initialized"
inherited_db.write_bytes(sentinel)
inherited_url = f"sqlite:///{inherited_db}"
truncation = "tests/test_truncate_message_count_regression.py"
owner = "tests/test_manage_tasks_owner_scope.py::test_edit_allowed_for_matching_owner"
manifest = [truncation, owner] if truncation_first else [owner, truncation]
script = textwrap.dedent('''
import os
import sys
import pytest
def snapshot():
import core
import src
import core.database as db
import core.session_manager as sm
import src.database as compat
return (
os.environ.get("DATABASE_URL"),
sys.modules["core.database"], core.database,
sys.modules["core.session_manager"], core.session_manager,
sys.modules["src.database"], src.database,
db.DATABASE_URL, db.engine, db.SessionLocal, db.Base,
db.Session, db.ChatMessage, db.ScheduledTask, db.set_sqlite_pragma,
compat.DATABASE_URL, compat.engine, compat.SessionLocal,
compat.Session, compat.ChatMessage,
sm.SessionLocal, sm.DbSession, sm.DbChatMessage, sm.SessionManager,
)
class StateGuard:
def pytest_sessionstart(self):
self.initial = snapshot()
assert self.initial[0] == "sqlite:///:memory:"
def pytest_collection_finish(self):
assert snapshot() == self.initial, "collection changed database bindings"
@pytest.hookimpl(hookwrapper=True, tryfirst=True)
def pytest_runtest_teardown(self):
yield
assert snapshot() == self.initial, "test leaked database or module state"
inherited_url = os.environ["DATABASE_URL"]
result = pytest.main(["-q", "-p", "no:cacheprovider", *sys.argv[1:]],
plugins=[StateGuard()])
assert os.environ["DATABASE_URL"] == inherited_url
raise SystemExit(result)
''')
result = subprocess.run(
[sys.executable, "-c", script, *manifest],
cwd=Path(__file__).resolve().parents[1],
env={**os.environ, "DATABASE_URL": inherited_url},
capture_output=True, text=True, timeout=60,
)
assert result.returncode == 0, result.stdout + result.stderr
assert "3 passed" in result.stdout
assert inherited_db.read_bytes() == sentinel
assert sorted(path.name for path in tmp_path.iterdir()) == ["developer.db"]
+9 -14
View File
@@ -5,36 +5,31 @@ document route tests. This keeps coverage on the real closures without spinning
up middleware.
"""
import tempfile
import uuid
import sys
from types import SimpleNamespace
from unittest.mock import MagicMock
import pytest
from fastapi import HTTPException
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
from sqlalchemy.pool import NullPool
from tests.helpers.database import disposable_database
from tests.helpers.import_state import clear_fake_database_modules
clear_fake_database_modules()
import core.database as cdb
import routes.document_routes as droutes
from core.database import Document
from core.database import Session as DbSession
from routes.document_helpers import DocumentPatch
from routes.document_helpers import _owner_session_filter
_TMPDB = tempfile.NamedTemporaryFile(suffix=".db", delete=False)
_ENGINE = create_engine(
f"sqlite:///{_TMPDB.name}",
connect_args={"check_same_thread": False},
poolclass=NullPool,
)
cdb.Base.metadata.create_all(_ENGINE)
_TS = sessionmaker(bind=_ENGINE, autoflush=False, autocommit=False)
@pytest.fixture(autouse=True)
def _document_database(tmp_path):
with disposable_database(tmp_path) as factory:
with pytest.MonkeyPatch.context() as patcher:
patcher.setattr(sys.modules[__name__], "_TS", factory, raising=False)
yield
def _req(user="alice"):
@@ -4,22 +4,22 @@ When AUTH_ENABLED=false, get_current_user returns None and gallery routes should
stay all-visible. When AUTH_ENABLED=true and no current user resolves, the same
None means an anonymous caller and gallery queries must fail closed.
"""
import tempfile
import uuid
import sys
import pytest
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
from sqlalchemy.pool import NullPool
from tests.helpers.database import disposable_database
import core.database as cdb
from core.database import GalleryImage
from routes.gallery_helpers import _owner_filter
_TMPDB = tempfile.NamedTemporaryFile(suffix=".db", delete=False)
_ENGINE = create_engine(f"sqlite:///{_TMPDB.name}", connect_args={"check_same_thread": False}, poolclass=NullPool)
cdb.Base.metadata.create_all(_ENGINE)
_TS = sessionmaker(bind=_ENGINE, autoflush=False, autocommit=False)
@pytest.fixture(autouse=True)
def _gallery_database(tmp_path):
with disposable_database(tmp_path) as factory:
with pytest.MonkeyPatch.context() as patcher:
patcher.setattr(sys.modules[__name__], "_TS", factory, raising=False)
yield
def _seed(*owners):
+12 -15
View File
@@ -12,14 +12,12 @@ permissive than the reader.
"""
import json
import tempfile
import sys
from datetime import datetime
import pytest
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
from sqlalchemy.pool import NullPool
from tests.helpers.database import disposable_database
from tests.helpers.import_state import clear_fake_database_modules
clear_fake_database_modules()
@@ -28,17 +26,16 @@ import core.database as cdb
from core.database import ScheduledTask
from src.tools.system import do_manage_tasks
_TMPDB = tempfile.NamedTemporaryFile(suffix=".db", delete=False)
_ENGINE = create_engine(
f"sqlite:///{_TMPDB.name}",
connect_args={"check_same_thread": False},
poolclass=NullPool,
)
cdb.Base.metadata.create_all(_ENGINE)
_TS = sessionmaker(bind=_ENGINE, autoflush=False, autocommit=False)
# do_manage_tasks does `from core.database import SessionLocal` at call time,
# so patching the module attribute is enough to point it at the temp DB.
cdb.SessionLocal = _TS
@pytest.fixture(autouse=True)
def _task_database(tmp_path):
# do_manage_tasks imports SessionLocal at call time. Own this binding for
# just one test, including helpers that seed and inspect its rows.
with disposable_database(tmp_path) as factory:
with pytest.MonkeyPatch.context() as patcher:
patcher.setattr(sys.modules[__name__], "_TS", factory, raising=False)
patcher.setattr(cdb, "SessionLocal", factory)
yield
def _seed(task_id, owner, *, name=None):
+11 -20
View File
@@ -9,30 +9,21 @@ inconsistent with the actual rows. get_session relies on message_count>0 to
decide whether to lazily hydrate from the DB, so an inflated count is a latent
correctness hazard.
"""
import os
import tempfile
import pytest
from tests.helpers.database import isolated_session_database
def _make_manager():
db_fd, db_path = tempfile.mkstemp(suffix=".db")
os.close(db_fd)
os.environ["DATABASE_URL"] = f"sqlite:///{db_path}"
# Import after DATABASE_URL is set so the engine binds to the temp DB.
import importlib
import core.database as database
importlib.reload(database)
database.Base.metadata.create_all(bind=database.engine)
import core.session_manager as sm_mod
importlib.reload(sm_mod)
return sm_mod.SessionManager(), database, sm_mod
@pytest.fixture
def manager_database(tmp_path):
with isolated_session_database(tmp_path) as resources:
yield resources
def test_truncate_keep_count_exceeds_total_does_not_inflate_count():
def test_truncate_keep_count_exceeds_total_does_not_inflate_count(manager_database):
from core.models import ChatMessage
sm, database, sm_mod = _make_manager()
sm, database = manager_database
sid = "short-session"
sm.create_session(session_id=sid, name="t", endpoint_url="x",
model="m", rag=False, owner="u")
@@ -59,10 +50,10 @@ def test_truncate_keep_count_exceeds_total_does_not_inflate_count():
db.close()
def test_truncate_keeps_history_alias_for_context_messages():
def test_truncate_keeps_history_alias_for_context_messages(manager_database):
from core.models import ChatMessage
sm, database, sm_mod = _make_manager()
sm, database = manager_database
sid = "alias-after-truncate"
sm.create_session(session_id=sid, name="t", endpoint_url="x",
model="m", rag=False, owner="u")