"""Shared data sources for contributor endpoints.""" import logging from contextlib import asynccontextmanager from typing import Any, AsyncIterator, Dict, Optional from fastapi import FastAPI from neo4j import Driver, GraphDatabase from owsclient import OwsClient from python_pdp_sdk.backends.authorization_backend import ( AuthorizationBackend, PdpAuthorizationBackend, ) from python_pdp_sdk.connectors.ows_pdp.ows_pdp import OwsPdpClient from contributor import config from contributor.api import context from contributor.connectors.features import SplitioClient, splitio_client_factory from contributor.connectors.redis import RedisConnector logger = logging.getLogger(__name__) AUTHORIZATION_BACKEND_KEY = "AUTHORIZATION_BACKEND" OWS_CLIENT_KEY = "OWS_CLIENT" NEO4J_DRIVER_KEY = "NEO4J_DRIVER" REDIS_CONNECTOR_KEY = "REDIS_CONNECTOR" SPLITIO_CLIENT_KEY = "SPLITIO_CLIENT" DATA_SOURCES: Dict[str, Any] = {} @asynccontextmanager async def datasources_lifespan( app: FastAPI | None, ) -> AsyncIterator[Dict[str, Any]]: ows_client = OwsClient( environment=config.ENVIRONMENT, service_name=config.SERVICE_NAME, correlation_id_getter=context.get_correlation_id, request_context_getter=context.get_request_context, ) DATA_SOURCES[OWS_CLIENT_KEY] = ows_client authorization_backend = setup_authorization_backend(ows_client) DATA_SOURCES[AUTHORIZATION_BACKEND_KEY] = authorization_backend splitio_client = splitio_client_factory() DATA_SOURCES[SPLITIO_CLIENT_KEY] = splitio_client logger.info("[lifespan] Initialized splitio client") if config.NEO4J_URL: neo4j_driver = GraphDatabase.driver( config.NEO4J_URL, auth=(config.NEO4J_USERNAME, config.NEO4J_PASSWORD), ) DATA_SOURCES[NEO4J_DRIVER_KEY] = neo4j_driver logger.info("[lifespan] Initialized Neo4j driver") redis_connector = RedisConnector( redis_host=config.REDIS_HOST, redis_port=config.REDIS_PORT, use_redis_cache=config.CACHE_USE_REDIS, ) DATA_SOURCES[REDIS_CONNECTOR_KEY] = redis_connector logger.info("[lifespan] Initialized Redis connector") try: yield DATA_SOURCES finally: if REDIS_CONNECTOR_KEY in DATA_SOURCES: await DATA_SOURCES[REDIS_CONNECTOR_KEY].close() if NEO4J_DRIVER_KEY in DATA_SOURCES: DATA_SOURCES[NEO4J_DRIVER_KEY].close() DATA_SOURCES.clear() def setup_authorization_backend(ows_client: OwsClient) -> AuthorizationBackend: """Set up the authorization backend for the application.""" ows_pdp_client = OwsPdpClient(ows_client) return PdpAuthorizationBackend(ows_pdp_client) def get_authorization_backend() -> AuthorizationBackend: """Dependency injector for the authorization backend.""" return DATA_SOURCES[AUTHORIZATION_BACKEND_KEY] def get_ows_client() -> OwsClient: """Dependency injector for the OWS client.""" return DATA_SOURCES[OWS_CLIENT_KEY] def get_neo4j_driver() -> Optional[Driver]: """Dependency injector for the Neo4j driver.""" return DATA_SOURCES.get(NEO4J_DRIVER_KEY) def get_redis_client() -> RedisConnector: """Dependency injector method for the redis connector.""" redis = DATA_SOURCES.get(REDIS_CONNECTOR_KEY) assert redis return redis def get_splitio_client() -> SplitioClient: """Dependency injector for the splitio client.""" return DATA_SOURCES[SPLITIO_CLIENT_KEY]