Files
odysseus/tests/test_database_test_isolation.py
T

136 lines
5.9 KiB
Python

"""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"]