diff --git a/wavefront/server/apps/floware/floware/server.py b/wavefront/server/apps/floware/floware/server.py index da58bd85..b881a493 100644 --- a/wavefront/server/apps/floware/floware/server.py +++ b/wavefront/server/apps/floware/floware/server.py @@ -26,7 +26,7 @@ from common_module.log.logger import logger from common_module.prometheus.prometheus_middleware import PrometheusMiddleware from common_module.response_formatter import ResponseFormatter -from db_repo_module.cache.azure_redis_auth import patch_redis_for_azure + from db_repo_module.database.connection import DatabaseClient from db_repo_module.db_repo_container import DatabaseModuleContainer from fastapi import HTTPException @@ -236,7 +236,6 @@ @asynccontextmanager async def lifespan(app: FastAPI): - patch_redis_for_azure() # Startup code (runs before the application starts) logger.info('Starting application...') diff --git a/wavefront/server/background_jobs/celery_worker/celery_worker/celery_app.py b/wavefront/server/background_jobs/celery_worker/celery_worker/celery_app.py index 81cf17f5..69e90a4f 100644 --- a/wavefront/server/background_jobs/celery_worker/celery_worker/celery_app.py +++ b/wavefront/server/background_jobs/celery_worker/celery_worker/celery_app.py @@ -1,6 +1,5 @@ from celery import Celery from celery.signals import ( - worker_process_init, worker_process_shutdown, worker_shutdown, ) @@ -8,13 +7,6 @@ from celery_worker.env import CELERY_BROKER_URL, CELERY_RESULT_BACKEND -@worker_process_init.connect -def setup_azure_redis_auth(**kwargs): - from db_repo_module.cache.azure_redis_auth import patch_redis_for_azure - - patch_redis_for_azure() - - def teardown_event_loop(**kwargs): from celery_worker.worker_setup import close_event_loop diff --git a/wavefront/server/modules/db_repo_module/db_repo_module/cache/azure_redis_auth.py b/wavefront/server/modules/db_repo_module/db_repo_module/cache/azure_redis_auth.py deleted file mode 100644 index c94180d7..00000000 --- a/wavefront/server/modules/db_repo_module/db_repo_module/cache/azure_redis_auth.py +++ /dev/null @@ -1,49 +0,0 @@ -import os -import threading - -import redis - -_patched = False - - -def patch_redis_for_azure() -> None: - global _patched - if _patched or os.getenv('CLOUD_PROVIDER', '').lower() != 'azure': - return - - from redis_entraid.cred_provider import create_from_default_azure_credential - from redis.credentials import CredentialProvider - - _inner = create_from_default_azure_credential(('https://redis.azure.com/.default',)) - - class _TimedProvider(CredentialProvider): - def get_credentials(self): - result = [] - exc = [] - - def _fetch(): - try: - result.append(_inner.get_credentials()) - except Exception as e: - exc.append(e) - - t = threading.Thread(target=_fetch, daemon=True) - t.start() - t.join(timeout=10) - if t.is_alive(): - raise TimeoutError('Azure Redis token fetch timed out after 10s') - if exc: - raise exc[0] - return result[0] - - provider = _TimedProvider() - original_init = redis.ConnectionPool.__init__ - - def patched_init(self, *args, **kw): - kw.pop('password', None) - kw.pop('username', None) - kw['credential_provider'] = provider - original_init(self, *args, **kw) - - redis.ConnectionPool.__init__ = patched_init - _patched = True