Chat-Service / app /memory /redis_checkpointer.py
ArabicNewsAnalyzer's picture
Update app/memory/redis_checkpointer.py
ead8f13 verified
Raw
History Blame Contribute Delete
2.72 kB
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