import contextlib import json import os from collections.abc import Iterator from pathlib import Path from typing import Any from unittest import mock import pytest from anydi import Container from confluent_kafka import Producer from fansifter_common.adapters.db.base import Database from fansifter_common.testing.factories import FactoryService from mypy_boto3_kms import KMSClient from app.adapters.db.orm import Model from app.container import container as _container from app.validator import RequestValidator from tests.unit.adapters.db.base import TestDatabase from tests.unit.types import CreateModel @pytest.fixture(scope="session") def container() -> Container: return _container @pytest.fixture(scope="session") def producer_mock() -> mock.MagicMock: return mock.MagicMock() @pytest.fixture(scope="session") def kms_client_mock() -> mock.MagicMock: return mock.MagicMock() @pytest.fixture(scope="session") def request_validator_mock() -> mock.MagicMock: return mock.MagicMock() @pytest.fixture(scope="function", autouse=True) def reset_mock(producer_mock: mock.MagicMock) -> None: producer_mock.reset_mock() @pytest.fixture(scope="session") def db() -> Iterator[Database]: # Create the database file before the tests db_path = Path(__file__).parent.parent.parent / "test.db" db = TestDatabase( url="sqlite:///" + str(db_path), session_args={ "expire_on_commit": False, "autoflush": True, }, ) Model.metadata.create_all(bind=db.engine) yield db db.close() Model.metadata.drop_all(bind=db.engine) # # Delete the database file after the tests db_path.unlink(missing_ok=True) @contextlib.contextmanager def _transaction_context(db: Database) -> Iterator[None]: with db.engine.begin() as conn, db.session_factory(bind=conn): try: yield finally: conn.rollback() @pytest.fixture(autouse=True) def _db_marker(request: pytest.FixtureRequest) -> Iterator[None]: marker = request.node.get_closest_marker("db") if marker is None: yield return db: Database = request.getfixturevalue("db") with _transaction_context(db): yield @pytest.fixture(scope="session") def factory_service() -> FactoryService: factory_service = FactoryService() return factory_service @pytest.fixture def create_model(factory_service: FactoryService, db: Database) -> CreateModel: def wrapper[T: Model](model: type[T], **kwargs: Any) -> T: return factory_service.create(db.session, model, **kwargs) return wrapper @pytest.fixture(scope="session", autouse=True) def override_test_dependencies( container: Container, producer_mock: mock.MagicMock, db: Database, kms_client_mock: mock.MagicMock, request_validator_mock: mock.MagicMock, ) -> None: """Override container dependencies.""" container.register( Producer, lambda: producer_mock, scope="singleton", override=True, ) container.register( Database, lambda: db, scope="singleton", override=True, ) container.register( KMSClient, lambda: kms_client_mock, scope="singleton", override=True, ) container.register( RequestValidator, lambda: request_validator_mock, scope="singleton", override=True, ) @pytest.fixture(scope="session") def test_event() -> dict[str, Any]: tests_path = os.path.dirname(os.path.abspath(__file__)) with open(f"{tests_path}/event.json") as fo: return json.load(fo)