from __future__ import annotations import asyncio import logging from redis.asyncio import Redis from redis.asyncio.retry import Retry from redis.backoff import ExponentialBackoff from redis.exceptions import ConnectionError, TimeoutError from langgraph.checkpoint.redis.aio import AsyncRedisSaver from app.config import get_settings logger = logging.getLogger(__name__) _checkpointer_cm = None _checkpointer: AsyncRedisSaver | None = None _redis_client: Redis | None = None _heartbeat_task: asyncio.Task | None = None def _build_redis_client(redis_url: str) -> Redis: return Redis.from_url( redis_url, health_check_interval=15, socket_keepalive=True, socket_connect_timeout=5, socket_timeout=10, retry_on_timeout=True, retry_on_error=[ConnectionError, TimeoutError], retry=Retry(ExponentialBackoff(base=0.5, cap=3), retries=3), ) async def _redis_heartbeat(interval: int = 15): """Keeps the pooled Redis connection alive by sending real traffic through it periodically — Railway's proxy kills idle connections, and OS-level TCP keepalive alone doesn't count as activity to it.""" while True: try: await asyncio.sleep(interval) if _redis_client is not None: await _redis_client.ping() except asyncio.CancelledError: break except Exception as e: logger.warning("Redis heartbeat ping failed: %s", e) async def get_checkpointer() -> AsyncRedisSaver: global _checkpointer_cm, _checkpointer, _redis_client, _heartbeat_task if _checkpointer is None: settings = get_settings() ttl_minutes = max(1, settings.chat_session_ttl_seconds // 60) _redis_client = _build_redis_client(settings.redis_url) _checkpointer_cm = AsyncRedisSaver.from_conn_string( redis_client=_redis_client, ttl={"default_ttl": ttl_minutes, "refresh_on_read": True}, ) _checkpointer = await _checkpointer_cm.__aenter__() await _checkpointer.asetup() _heartbeat_task = asyncio.create_task(_redis_heartbeat()) return _checkpointer async def close_checkpointer() -> None: global _checkpointer_cm, _checkpointer, _redis_client, _heartbeat_task if _heartbeat_task is not None: _heartbeat_task.cancel() if _checkpointer_cm is not None: await _checkpointer_cm.__aexit__(None, None, None) if _redis_client is not None: await _redis_client.aclose() _checkpointer = None _checkpointer_cm = None _redis_client = None _heartbeat_task = None