Skip to content

Commit 248c061

Browse files
authored
fix(model-engine): harden aioredis pool against stale-socket errors (#816)
1 parent 237b404 commit 248c061

10 files changed

Lines changed: 117 additions & 9 deletions

File tree

model-engine/model_engine_server/api/dependencies.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,7 @@
77
import redis.asyncio as aioredis
88
from fastapi import Depends, HTTPException, status
99
from fastapi.security import HTTPBasic, HTTPBasicCredentials, OAuth2PasswordBearer
10+
from model_engine_server.common.aioredis_pool import build_aioredis_pool
1011
from model_engine_server.common.config import hmi_config
1112
from model_engine_server.common.dtos.model_endpoints import BrokerType
1213
from model_engine_server.common.env_vars import CIRCLECI
@@ -560,6 +561,6 @@ def get_or_create_aioredis_pool() -> aioredis.ConnectionPool:
560561

561562
expiration_timestamp = hmi_config.cache_redis_url_expiration_timestamp
562563
if _pool is None or (expiration_timestamp is not None and time.time() > expiration_timestamp):
563-
_pool = aioredis.BlockingConnectionPool.from_url(hmi_config.cache_redis_url)
564+
_pool = build_aioredis_pool(hmi_config.cache_redis_url)
564565
assert _pool is not None
565566
return _pool
Lines changed: 52 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,52 @@
1+
"""Shared construction for aioredis connection pools and clients.
2+
3+
redis-py defaults leave pooled connections vulnerable to middlebox idle
4+
timeouts (Istio sidecars, cloud load balancers, NAT): when `retry` and
5+
`retry_on_error` are both unset, `AbstractConnection.__init__` uses
6+
`Retry(NoBackoff(), 0)` — zero retries — so the first command issued on a
7+
silently-closed pooled socket surfaces as `ConnectionError` to the caller.
8+
9+
The helpers here turn on TCP keepalive, a pre-command PING after idle periods,
10+
and transparent retry with exponential backoff. Use them for every long-lived
11+
aioredis pool or client in this repo.
12+
"""
13+
14+
from typing import Any, Dict
15+
16+
import redis.asyncio as aioredis
17+
from redis.asyncio.retry import Retry
18+
from redis.backoff import ExponentialBackoff
19+
from redis.exceptions import ConnectionError as RedisConnectionError
20+
from redis.exceptions import TimeoutError as RedisTimeoutError
21+
22+
HEALTH_CHECK_INTERVAL_SECONDS = 30
23+
SOCKET_CONNECT_TIMEOUT_SECONDS = 5
24+
RETRY_BACKOFF_CAP_SECONDS = 2.0
25+
RETRY_BACKOFF_BASE_SECONDS = 0.2
26+
RETRY_MAX_ATTEMPTS = 3
27+
28+
29+
def get_aioredis_connection_kwargs() -> Dict[str, Any]:
30+
"""Return kwargs for constructing a resilient aioredis pool or client.
31+
32+
A fresh Retry and error list are returned on every call so callers can
33+
never accidentally share mutable state across pools.
34+
"""
35+
return {
36+
"socket_keepalive": True,
37+
"socket_connect_timeout": SOCKET_CONNECT_TIMEOUT_SECONDS,
38+
"health_check_interval": HEALTH_CHECK_INTERVAL_SECONDS,
39+
"retry_on_error": [RedisConnectionError, RedisTimeoutError],
40+
"retry": Retry(
41+
ExponentialBackoff(cap=RETRY_BACKOFF_CAP_SECONDS, base=RETRY_BACKOFF_BASE_SECONDS),
42+
retries=RETRY_MAX_ATTEMPTS,
43+
),
44+
}
45+
46+
47+
def build_aioredis_pool(url: str) -> aioredis.ConnectionPool:
48+
return aioredis.BlockingConnectionPool.from_url(url, **get_aioredis_connection_kwargs())
49+
50+
51+
def build_aioredis_client(url: str) -> aioredis.Redis:
52+
return aioredis.from_url(url, **get_aioredis_connection_kwargs())

model-engine/model_engine_server/core/celery/app.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -9,6 +9,7 @@
99
from celery.app import backends
1010
from celery.app.control import Inspect
1111
from celery.result import AsyncResult
12+
from model_engine_server.common.aioredis_pool import build_aioredis_client
1213
from model_engine_server.core.aws.roles import session
1314
from model_engine_server.core.aws.secrets import get_key_file
1415
from model_engine_server.core.config import infra_config
@@ -232,7 +233,7 @@ def get_redis_instance(db_index: int = 0) -> Union[Redis, StrictRedis]:
232233

233234
def get_async_redis_instance(db_index: int = 0) -> aioredis.Redis:
234235
host, port = get_redis_host_port()
235-
return aioredis.Redis.from_url(f"redis://{host}:{port}/{db_index}")
236+
return build_aioredis_client(f"redis://{host}:{port}/{db_index}")
236237

237238

238239
def celery_app(

model-engine/model_engine_server/core/celery/celery_autoscaler.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -11,7 +11,6 @@
1111
from math import ceil
1212
from typing import Any, DefaultDict, Dict, List, Set, Tuple
1313

14-
import redis.asyncio as aioredis
1514
import stringcase
1615
from azure.core.exceptions import ResourceNotFoundError
1716
from azure.identity import DefaultAzureCredential
@@ -22,6 +21,7 @@
2221
from kubernetes_asyncio import config as kube_config
2322
from kubernetes_asyncio.client.rest import ApiException
2423
from kubernetes_asyncio.config.config_exception import ConfigException
24+
from model_engine_server.common.aioredis_pool import build_aioredis_client
2525
from model_engine_server.core.aws.roles import session
2626
from model_engine_server.core.celery import (
2727
TaskVisibility,
@@ -350,7 +350,7 @@ async def _init_client(self):
350350
get_redis_host_port()
351351
) # Switches the redis instance based on CELERY_ELASTICACHE_ENABLED's value
352352
self.redis = {
353-
db_index: aioredis.Redis.from_url(f"redis://{host}:{port}/{db_index}")
353+
db_index: build_aioredis_client(f"redis://{host}:{port}/{db_index}")
354354
for db_index in get_all_db_indexes()
355355
}
356356
self.initialized = True

model-engine/model_engine_server/entrypoints/start_batch_job_orchestration.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,7 @@
66

77
import redis.asyncio as aioredis
88
from model_engine_server.api.dependencies import get_monitoring_metrics_gateway
9+
from model_engine_server.common.aioredis_pool import build_aioredis_pool
910
from model_engine_server.common.config import hmi_config
1011
from model_engine_server.common.dtos.model_endpoints import BrokerType
1112
from model_engine_server.common.env_vars import CIRCLECI
@@ -65,7 +66,7 @@ async def run_batch_job(
6566
):
6667
tracing_gateway = get_tracing_gateway()
6768
session = get_session_async_null_pool()
68-
pool = aioredis.BlockingConnectionPool.from_url(hmi_config.cache_redis_url)
69+
pool = build_aioredis_pool(hmi_config.cache_redis_url)
6970
redis: aioredis.Redis[Any] = aioredis.Redis(connection_pool=pool)
7071
sqs_task_queue_gateway = CeleryTaskQueueGateway(
7172
broker_type=BrokerType.SQS, tracing_gateway=tracing_gateway

model-engine/model_engine_server/infra/gateways/redis_inference_autoscaling_metrics_gateway.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
11
from typing import Optional
22

33
import redis.asyncio as aioredis
4+
from model_engine_server.common.aioredis_pool import build_aioredis_client
45
from model_engine_server.domain.gateways.inference_autoscaling_metrics_gateway import (
56
InferenceAutoscalingMetricsGateway,
67
)
@@ -19,7 +20,7 @@ def __init__(
1920
# If aioredis cannot create a connection pool, reraise that as an error because the
2021
# default error message is cryptic and not obvious.
2122
try:
22-
self._redis = aioredis.from_url(redis_info, health_check_interval=60)
23+
self._redis = build_aioredis_client(redis_info)
2324
except Exception as exc:
2425
raise RuntimeError(
2526
"If redis_info is specified, RedisInferenceAutoscalingMetricsGateway must be"

model-engine/model_engine_server/infra/repositories/redis_feature_flag_repository.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
11
from typing import Optional
22

33
import redis.asyncio as aioredis
4+
from model_engine_server.common.aioredis_pool import build_aioredis_client
45
from model_engine_server.infra.repositories.feature_flag_repository import FeatureFlagRepository
56

67

@@ -15,7 +16,7 @@ def __init__(
1516
# If aioredis cannot create a connection pool, reraise that as an error because the
1617
# default error message is cryptic and not obvious.
1718
try:
18-
self._redis = aioredis.from_url(redis_info, health_check_interval=60)
19+
self._redis = build_aioredis_client(redis_info)
1920
except Exception as exc:
2021
raise RuntimeError(
2122
"If redis_info is specified, RedisFeatureFlagRepository must be"

model-engine/model_engine_server/infra/repositories/redis_model_endpoint_cache_repository.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,7 @@
33
from typing import Optional
44

55
import redis.asyncio as aioredis
6+
from model_engine_server.common.aioredis_pool import build_aioredis_client
67
from model_engine_server.domain.entities import ModelEndpointInfraState
78
from model_engine_server.infra.repositories.model_endpoint_cache_repository import (
89
ModelEndpointCacheRepository,
@@ -23,7 +24,7 @@ def __init__(
2324
# If aioredis cannot create a connection pool, reraise that as an error because the
2425
# default error message is cryptic and not obvious.
2526
try:
26-
self._redis = aioredis.from_url(redis_info, health_check_interval=60)
27+
self._redis = build_aioredis_client(redis_info)
2728
except Exception as exc:
2829
raise RuntimeError(
2930
"If redis_info is specified, RedisModelEndpointCacheRepository must be"

model-engine/model_engine_server/service_builder/tasks_v1.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,7 @@
88
from celery.signals import worker_process_init
99
from celery.utils.log import get_task_logger
1010
from model_engine_server.api.dependencies import get_monitoring_metrics_gateway
11+
from model_engine_server.common.aioredis_pool import build_aioredis_pool
1112
from model_engine_server.common.config import hmi_config
1213
from model_engine_server.common.constants import READYZ_FPATH
1314
from model_engine_server.common.dtos.endpoint_builder import (
@@ -155,7 +156,7 @@ async def _build_endpoint(
155156
# Redis connection
156157
if infra_config().debug_mode: # pragma: no cover
157158
logger.info("Connecting to Redis", extra={"redis_url": hmi_config.cache_redis_url})
158-
pool = aioredis.BlockingConnectionPool.from_url(hmi_config.cache_redis_url)
159+
pool = build_aioredis_pool(hmi_config.cache_redis_url)
159160
redis = aioredis.Redis(connection_pool=pool)
160161
if infra_config().debug_mode: # pragma: no cover
161162
logger.info("Redis connection established successfully")
Lines changed: 49 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,49 @@
1+
import redis.asyncio as aioredis
2+
from model_engine_server.common.aioredis_pool import (
3+
HEALTH_CHECK_INTERVAL_SECONDS,
4+
RETRY_MAX_ATTEMPTS,
5+
SOCKET_CONNECT_TIMEOUT_SECONDS,
6+
build_aioredis_client,
7+
build_aioredis_pool,
8+
get_aioredis_connection_kwargs,
9+
)
10+
from redis.asyncio.retry import Retry
11+
from redis.exceptions import ConnectionError as RedisConnectionError
12+
from redis.exceptions import TimeoutError as RedisTimeoutError
13+
14+
15+
def test_kwargs_include_keepalive_health_check_and_retry():
16+
kwargs = get_aioredis_connection_kwargs()
17+
assert kwargs["socket_keepalive"] is True
18+
assert kwargs["socket_connect_timeout"] == SOCKET_CONNECT_TIMEOUT_SECONDS
19+
assert kwargs["health_check_interval"] == HEALTH_CHECK_INTERVAL_SECONDS
20+
assert kwargs["retry_on_error"] == [RedisConnectionError, RedisTimeoutError]
21+
assert isinstance(kwargs["retry"], Retry)
22+
23+
24+
def test_kwargs_do_not_share_mutable_state_across_calls():
25+
first = get_aioredis_connection_kwargs()
26+
second = get_aioredis_connection_kwargs()
27+
assert first["retry"] is not second["retry"]
28+
assert first["retry_on_error"] is not second["retry_on_error"]
29+
30+
31+
def test_build_aioredis_pool_applies_kwargs_to_connection():
32+
pool = build_aioredis_pool("redis://localhost:6379/0")
33+
assert isinstance(pool, aioredis.BlockingConnectionPool)
34+
conn = pool.connection_class(**pool.connection_kwargs)
35+
assert conn.socket_keepalive is True
36+
assert conn.socket_connect_timeout == SOCKET_CONNECT_TIMEOUT_SECONDS
37+
assert conn.health_check_interval == HEALTH_CHECK_INTERVAL_SECONDS
38+
# retries is the only reliable way to tell we replaced the default
39+
# NoBackoff(0) with our configured Retry.
40+
assert conn.retry._retries == RETRY_MAX_ATTEMPTS
41+
42+
43+
def test_build_aioredis_client_applies_kwargs_to_connection():
44+
client = build_aioredis_client("redis://localhost:6379/0")
45+
pool = client.connection_pool
46+
conn = pool.connection_class(**pool.connection_kwargs)
47+
assert conn.socket_keepalive is True
48+
assert conn.health_check_interval == HEALTH_CHECK_INTERVAL_SECONDS
49+
assert conn.retry._retries == RETRY_MAX_ATTEMPTS

0 commit comments

Comments
 (0)