mirror of
https://github.com/pewdiepie-archdaemon/odysseus.git
synced 2026-10-06 06:52:20 +02:00
fix(tests): isolate database and module state
This commit is contained in:
+8
-9
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -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():
|
||||
|
||||
@@ -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"]
|
||||
@@ -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,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):
|
||||
|
||||
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user