from unittest import mock import sqlalchemy as sa from anydi import Module, provider from jinja2sql import Jinja2SQL from snowflake.snowpark import Session as SnowparkSession from dmp.adapters.db import DefaultDB, ReportingDB from dmp.config import Settings from tests.unit.adapters.db import TestDefaultDB, TestReportingDB class TestModule(Module): @provider(scope="singleton", override=True) def pg_db(self, settings: Settings, jinja2sql: Jinja2SQL) -> DefaultDB: return TestDefaultDB( url=settings.postgres_url, session_args={ "expire_on_commit": False, "autoflush": True, }, jinja2sql=jinja2sql, ) @provider(scope="singleton", override=True) def reporting_db(self, settings: Settings, jinja2sql: Jinja2SQL) -> ReportingDB: return TestReportingDB( url=sa.URL.create( drivername="postgresql+psycopg", username=settings.snowflake_user, password=settings.snowflake_password, host=settings.snowflake_host, port=settings.snowflake_port, database=settings.snowflake_database, ), session_args={ "expire_on_commit": False, "autoflush": True, }, jinja2sql=jinja2sql, ) @provider(scope="singleton", override=True) def llm_session(self) -> SnowparkSession: return mock.MagicMock(spec=SnowparkSession)