from collections.abc import Iterator import boto3 import redis import sqlalchemy.pool from anydi import Container, Module, Provider, provider from fansifter_common import context from fansifter_common.adapters.aws.secretsmanager import SecretsManager from fansifter_common.constants import PROD_ENVIRONMENT, QA_ENVIRONMENT from fansifter_common.m2m_token import M2MTokenManager from fansifter_common.utils.functional import lazy_proxy from jinja2sql import Jinja2SQL from owsclient import OwsClient from app.adapters.aws import LambdaClient from app.adapters.db import Database from app.adapters.ows_dmp import OwsDmpClient from app.config import Settings, settings from app.strategies import ( FairDistributionStrategy, FcfsStrategy, QuotaAllocationStrategy, ) class AppModule(Module): @provider(scope="singleton") def aws_session(self, settings: Settings) -> boto3.Session: return boto3.Session(region_name=settings.aws_region_name) @provider(scope="singleton") def lambda_client(self, settings: Settings, session: boto3.Session) -> LambdaClient: return LambdaClient( region_name=settings.aws_region_name, session=session, ) @provider(scope="singleton") def secrets_manager(self, settings: Settings) -> SecretsManager: return SecretsManager(region_name=settings.aws_region_name) @provider(scope="singleton") def m2m_token_manager( self, settings: Settings, secrets_manager: SecretsManager ) -> M2MTokenManager: return M2MTokenManager( secrets_manager=secrets_manager, secret_name_key=settings.m2m_token_secret_key_name, secret_expire_name_key=settings.m2m_token_secret_expiry_key_name, ) @provider(scope="singleton") def ows_client( self, settings: Settings, m2m_token_manager: M2MTokenManager ) -> OwsClient: return OwsClient( environment=settings.environment, service_name=settings.service_name, m2m_token_manager=m2m_token_manager if settings.environment in [QA_ENVIRONMENT, PROD_ENVIRONMENT] else None, correlation_id_getter=context.get_correlation_id, request_context_getter=context.get_request_context, ) @provider(scope="singleton") def ows_dmp_client(self, ows_client: OwsClient) -> OwsDmpClient: return OwsDmpClient(ows_client=ows_client) @provider(scope="singleton") def jinja2sql(self, settings: Settings) -> Jinja2SQL: jinja2sql = Jinja2SQL(searchpath=settings.jinja2sql_template_searchpath) jinja2sql.env.globals.update({"settings": settings}) # type: ignore return jinja2sql @provider(scope="singleton") def db(self, settings: Settings, jinja2sql: Jinja2SQL) -> Iterator[Database]: with Database( url=settings.snowflake_url, engine_args={ "echo": settings.snowflake_echo, "poolclass": sqlalchemy.pool.QueuePool, "pool_size": settings.snowflake_pool_size, "max_overflow": settings.snowflake_pool_max_overflow, "pool_recycle": settings.snowflake_pool_recycle, "pool_pre_ping": settings.snowflake_pool_pre_ping, "pool_reset_on_return": settings.snowflake_pool_reset_on_return, "connect_args": settings.snowflake_connect_args, }, jinja2sql=jinja2sql, ) as db: yield db @provider(scope="request") def db_session_factory(self, db: Database) -> Iterator[None]: with db.session_factory(): yield @provider(scope="singleton") def redis_client(self, settings: Settings) -> redis.Redis: return redis.Redis( host=settings.redis_host, port=settings.redis_port, db=settings.redis_db, ssl=settings.redis_ssl, ) @provider(scope="singleton") def quota_strategy( self, settings: Settings, ) -> QuotaAllocationStrategy: if settings.quota_strategy == "fcfs": return FcfsStrategy() return FairDistributionStrategy() def make_container() -> Container: return Container( providers=[ Provider(Settings, factory=lambda: settings, scope="singleton"), ], modules=[AppModule], ) container = lazy_proxy(make_container)