diff --git a/tests/conftest.py b/tests/conftest.py index 0915d519c..fae910d71 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -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 diff --git a/tests/helpers/database.py b/tests/helpers/database.py new file mode 100644 index 000000000..57c19dd69 --- /dev/null +++ b/tests/helpers/database.py @@ -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 diff --git a/tests/test_checkin_digest_owner_scope.py b/tests/test_checkin_digest_owner_scope.py index a2e8ebb17..68f1575d4 100644 --- a/tests/test_checkin_digest_owner_scope.py +++ b/tests/test_checkin_digest_owner_scope.py @@ -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(): diff --git a/tests/test_database_test_isolation.py b/tests/test_database_test_isolation.py new file mode 100644 index 000000000..f0e72e27a --- /dev/null +++ b/tests/test_database_test_isolation.py @@ -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"] diff --git a/tests/test_document_session_owner_scope.py b/tests/test_document_session_owner_scope.py index f776d9822..372092306 100644 --- a/tests/test_document_session_owner_scope.py +++ b/tests/test_document_session_owner_scope.py @@ -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"): diff --git a/tests/test_gallery_owner_filter_single_user.py b/tests/test_gallery_owner_filter_single_user.py index 7032410c6..215bc5911 100644 --- a/tests/test_gallery_owner_filter_single_user.py +++ b/tests/test_gallery_owner_filter_single_user.py @@ -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): diff --git a/tests/test_manage_tasks_owner_scope.py b/tests/test_manage_tasks_owner_scope.py index 14797a2f4..2ef5e7410 100644 --- a/tests/test_manage_tasks_owner_scope.py +++ b/tests/test_manage_tasks_owner_scope.py @@ -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): diff --git a/tests/test_truncate_message_count_regression.py b/tests/test_truncate_message_count_regression.py index 6f3d4ba0f..a632f8693 100644 --- a/tests/test_truncate_message_count_regression.py +++ b/tests/test_truncate_message_count_regression.py @@ -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")