from __future__ import annotations from collections.abc import Iterator import pytest import sqlalchemy as sa from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column from fansifter_common.adapters.db import Database class Base(DeclarativeBase): pass class Item(Base): __tablename__ = "item" id: Mapped[int] = mapped_column(primary_key=True) name: Mapped[str] @pytest.fixture() def db() -> Iterator[Database]: database = Database(url="sqlite:///:memory:") Base.metadata.create_all(database.engine) yield database database.close() def count_items(db: Database) -> int: with db.session_factory(): return db.session.execute( sa.select(sa.func.count()).select_from(Item) ).scalar_one() # --------------------------------------------------------------------------- # has_session / session # --------------------------------------------------------------------------- class TestSession: def test_has_session_false_without_session(self, db: Database) -> None: assert db.has_session() is False def test_has_session_true_inside_session_factory(self, db: Database) -> None: with db.session_factory(): assert db.has_session() is True def test_has_session_false_after_session_factory(self, db: Database) -> None: with db.session_factory(): pass assert db.has_session() is False def test_session_raises_without_session(self, db: Database) -> None: with pytest.raises(RuntimeError, match="Session is not started"): _ = db.session def test_session_returns_same_object_inside_session_factory( self, db: Database ) -> None: with db.session_factory() as session: assert db.session is session def test_session_factory_reuses_existing_session(self, db: Database) -> None: with db.session_factory() as outer: with db.session_factory() as inner: assert outer is inner def test_session_factory_replace_creates_new_session(self, db: Database) -> None: with db.session_factory() as outer: with db.session_factory(replace=True) as inner: assert outer is not inner # --------------------------------------------------------------------------- # in_transaction # --------------------------------------------------------------------------- class TestInTransaction: def test_false_without_any_context(self, db: Database) -> None: assert db.in_transaction() is False def test_false_inside_bare_session_factory(self, db: Database) -> None: with db.session_factory(): assert db.in_transaction() is False def test_true_inside_transaction(self, db: Database) -> None: with db.transaction(): assert db.in_transaction() is True def test_false_after_transaction(self, db: Database) -> None: with db.transaction(): pass assert db.in_transaction() is False def test_true_inside_rollback_transaction(self, db: Database) -> None: with db.rollback_transaction(): assert db.in_transaction() is True def test_false_after_rollback_transaction(self, db: Database) -> None: with db.rollback_transaction(): pass assert db.in_transaction() is False # --------------------------------------------------------------------------- # transaction — context manager # --------------------------------------------------------------------------- class TestTransactionContextManager: def test_commits_on_success(self, db: Database) -> None: with db.transaction(): db.session.add(Item(name="alice")) assert count_items(db) == 1 def test_rollback_on_exception(self, db: Database) -> None: with pytest.raises(ValueError): with db.transaction(): db.session.add(Item(name="alice")) raise ValueError assert count_items(db) == 0 def test_commit_on_error_matching_exception(self, db: Database) -> None: with pytest.raises(ValueError): with db.transaction(commit_on_error=ValueError): db.session.add(Item(name="alice")) raise ValueError assert count_items(db) == 1 def test_commit_on_error_non_matching_exception(self, db: Database) -> None: with pytest.raises(ValueError): with db.transaction(commit_on_error=RuntimeError): db.session.add(Item(name="alice")) raise ValueError assert count_items(db) == 0 def test_commit_on_error_tuple(self, db: Database) -> None: with pytest.raises(ValueError): with db.transaction(commit_on_error=(RuntimeError, ValueError)): db.session.add(Item(name="alice")) raise ValueError assert count_items(db) == 1 def test_nested_transaction_raises(self, db: Database) -> None: with db.transaction(): with pytest.raises(RuntimeError, match="Transaction already started"): with db.transaction(): pass # --------------------------------------------------------------------------- # transaction — decorator # --------------------------------------------------------------------------- class TestTransactionDecorator: def test_decorator_without_parens_commits(self, db: Database) -> None: @db.transaction def create() -> None: db.session.add(Item(name="alice")) create() assert count_items(db) == 1 def test_decorator_with_parens_commits(self, db: Database) -> None: @db.transaction() def create() -> None: db.session.add(Item(name="alice")) create() assert count_items(db) == 1 def test_decorator_with_commit_on_error(self, db: Database) -> None: @db.transaction(commit_on_error=ValueError) def create_and_fail() -> None: db.session.add(Item(name="alice")) raise ValueError with pytest.raises(ValueError): create_and_fail() assert count_items(db) == 1 def test_decorator_without_parens_rollback_on_exception(self, db: Database) -> None: @db.transaction def create_and_fail() -> None: db.session.add(Item(name="alice")) raise ValueError with pytest.raises(ValueError): create_and_fail() assert count_items(db) == 0 def test_decorator_preserves_return_value(self, db: Database) -> None: @db.transaction def create(name: str) -> str: db.session.add(Item(name=name)) return name result = create("alice") assert result == "alice" def test_decorator_preserves_signature(self, db: Database) -> None: """ParamSpec: type checker sees the original signature, not (*args, **kwargs).""" @db.transaction def create(name: str, *, active: bool = True) -> Item: item = Item(name=name) db.session.add(item) return item item = create("alice", active=False) assert item.name == "alice" # --------------------------------------------------------------------------- # autocommit — context manager # --------------------------------------------------------------------------- class TestAutocommitContextManager: def test_data_visible_after_block(self, db: Database) -> None: with db.autocommit(): db.session.execute(sa.insert(Item).values(name="alice")) assert count_items(db) == 1 def test_raises_inside_transaction(self, db: Database) -> None: with db.transaction(): with pytest.raises(RuntimeError, match="Cannot use autocommit"): with db.autocommit(): pass # --------------------------------------------------------------------------- # autocommit — decorator # --------------------------------------------------------------------------- class TestAutocommitDecorator: def test_decorator_without_parens(self, db: Database) -> None: @db.autocommit def create() -> None: db.session.execute(sa.insert(Item).values(name="alice")) create() assert count_items(db) == 1 def test_decorator_with_parens(self, db: Database) -> None: @db.autocommit() def create() -> None: db.session.execute(sa.insert(Item).values(name="alice")) create() assert count_items(db) == 1 def test_decorator_preserves_return_value(self, db: Database) -> None: @db.autocommit def create(name: str) -> str: db.session.execute(sa.insert(Item).values(name=name)) return name result = create("alice") assert result == "alice" # --------------------------------------------------------------------------- # rollback_transaction # --------------------------------------------------------------------------- class TestRollbackTransaction: def test_rolls_back_all_changes(self, db: Database) -> None: with db.rollback_transaction(): with db.transaction(): db.session.add(Item(name="alice")) assert count_items(db) == 0 def test_raises_when_session_already_exists(self, db: Database) -> None: with db.session_factory(): with pytest.raises(RuntimeError, match="Cannot use rollback_transaction"): with db.rollback_transaction(): pass def test_nested_transaction_data_visible_within_session(self, db: Database) -> None: """transaction() inside rollback_transaction() flushes (not commits), so data is visible within the same session but rolled back at the end.""" with db.rollback_transaction(): with db.transaction(): db.session.add(Item(name="alice")) # Inner transaction flushed — data visible within the same session. count = db.session.execute( sa.select(sa.func.count()).select_from(Item) ).scalar_one() assert count == 1 # Outer rollback undoes everything. assert count_items(db) == 0 def test_exception_inside_still_rolls_back(self, db: Database) -> None: with pytest.raises(ValueError): with db.rollback_transaction(): with db.transaction(): db.session.add(Item(name="alice")) raise ValueError assert count_items(db) == 0 def test_in_transaction_true_prevents_repo_commits(self, db: Database) -> None: """in_transaction() must be True inside rollback_transaction so that repository _commit_or_flush defers to flush, not commit.""" with db.rollback_transaction(): assert db.in_transaction() is True def test_autocommit_inside_rollback_transaction_rolls_back( self, db: Database ) -> None: """autocommit() inside rollback_transaction() must not raise and must still roll back — so test isolation is preserved even for code that uses autocommit internally.""" with db.rollback_transaction(): with db.autocommit(): db.session.execute(sa.insert(Item).values(name="alice")) assert count_items(db) == 0 # --------------------------------------------------------------------------- # global_context # --------------------------------------------------------------------------- class TestGlobalContext: def test_has_session_visible_from_spawned_thread(self, db: Database) -> None: """Inside global_context, session state set in one thread is visible in another.""" import threading seen_in_thread: list[bool] = [] with db.global_context(): with db.session_factory(): assert db.has_session() is True def check() -> None: seen_in_thread.append(db.has_session()) thread = threading.Thread(target=check) thread.start() thread.join() assert seen_in_thread == [True] def test_in_transaction_visible_from_spawned_thread(self, db: Database) -> None: """Transaction state set in one thread is visible in another under global_context.""" import threading seen_in_thread: list[bool] = [] with db.global_context(): with db.transaction(): assert db.in_transaction() is True def check() -> None: seen_in_thread.append(db.in_transaction()) thread = threading.Thread(target=check) thread.start() thread.join() assert seen_in_thread == [True] def test_context_vars_restored_after_exit(self, db: Database) -> None: """After global_context exits, ContextVars are back in effect.""" with db.global_context(): pass # Should behave normally — ContextVar-isolated, not visible across threads. assert db.has_session() is False assert db.in_transaction() is False def test_rollback_transaction_inside_global_context(self, db: Database) -> None: """global_context + rollback_transaction work together for test isolation.""" with db.global_context(): with db.rollback_transaction(): with db.transaction(): db.session.add(Item(name="alice")) assert count_items(db) == 0