Files
leocrm/tests/test_resilience.py
T

391 lines
12 KiB
Python
Raw Normal View History

"""Tests for resilience patterns: circuit breaker, DB retry, Redis fallback, in-memory rate limiter."""
import asyncio
import time
import pytest
from unittest.mock import AsyncMock, MagicMock, patch
from app.core.resilience import (
CircuitBreaker,
CircuitBreakerMiddleware,
InMemoryRateLimiter,
RedisUnavailableError,
get_circuit,
get_inmemory_limiter,
reset_all_circuits,
retry_db,
redis_call_with_fallback,
)
class TestCircuitBreaker:
"""Unit tests for CircuitBreaker."""
def setup_method(self):
self.cb = CircuitBreaker("test", failure_threshold=3, window_seconds=30, cooldown_seconds=60)
@pytest.mark.asyncio
async def test_closed_allows_requests(self):
assert await self.cb.can_proceed() is True
@pytest.mark.asyncio
async def test_opens_after_threshold(self):
for _ in range(3):
await self.cb.record_failure()
assert self.cb.state == CircuitBreaker.OPEN
assert await self.cb.can_proceed() is False
@pytest.mark.asyncio
async def test_success_resets_circuit(self):
for _ in range(2):
await self.cb.record_failure()
await self.cb.record_success()
assert self.cb.state == CircuitBreaker.CLOSED
assert len(self.cb._failures) == 0
@pytest.mark.asyncio
async def test_half_open_after_cooldown(self):
for _ in range(3):
await self.cb.record_failure()
assert self.cb.state == CircuitBreaker.OPEN
# Simulate cooldown expiry
self.cb._opened_at = time.monotonic() - 61
assert self.cb.state == CircuitBreaker.HALF_OPEN
# One probe allowed
assert await self.cb.can_proceed() is True
# Second probe blocked
assert await self.cb.can_proceed() is False
@pytest.mark.asyncio
async def test_half_open_probe_success_closes_circuit(self):
for _ in range(3):
await self.cb.record_failure()
self.cb._opened_at = time.monotonic() - 61
# Allow probe
await self.cb.can_proceed()
await self.cb.record_success()
assert self.cb.state == CircuitBreaker.CLOSED
@pytest.mark.asyncio
async def test_half_open_probe_failure_reopens(self):
for _ in range(3):
await self.cb.record_failure()
self.cb._opened_at = time.monotonic() - 61
await self.cb.can_proceed()
await self.cb.record_failure()
assert self.cb.state == CircuitBreaker.OPEN
@pytest.mark.asyncio
async def test_failures_outside_window_are_pruned(self):
cb = CircuitBreaker("test2", failure_threshold=2, window_seconds=1, cooldown_seconds=60)
await cb.record_failure()
# Wait for window to expire
await asyncio.sleep(1.1)
await cb.record_failure()
# Only 1 failure in current window — circuit should stay closed
assert cb.state == CircuitBreaker.CLOSED
def test_reset(self):
self.cb._state = CircuitBreaker.OPEN
self.cb._failures = [1.0, 2.0]
self.cb.reset()
assert self.cb.state == CircuitBreaker.CLOSED
assert len(self.cb._failures) == 0
class TestGetCircuit:
"""Test global circuit registry."""
def test_returns_same_instance(self):
cb1 = get_circuit("shared")
cb2 = get_circuit("shared")
assert cb1 is cb2
def test_different_names_different_instances(self):
cb1 = get_circuit("db")
cb2 = get_circuit("redis")
assert cb1 is not cb2
def test_reset_all(self):
get_circuit("db")._state = CircuitBreaker.OPEN
get_circuit("redis")._state = CircuitBreaker.OPEN
reset_all_circuits()
assert get_circuit("db").state == CircuitBreaker.CLOSED
assert get_circuit("redis").state == CircuitBreaker.CLOSED
class TestRetryDB:
"""Tests for DB retry decorator."""
@pytest.mark.asyncio
async def test_success_no_retry(self):
call_count = 0
async def success_op():
nonlocal call_count
call_count += 1
return "ok"
result = await retry_db(success_op, max_retries=3, base_delay=0.01)
assert result == "ok"
assert call_count == 1
@pytest.mark.asyncio
async def test_retries_on_transient_error(self):
call_count = 0
async def flaky_op():
nonlocal call_count
call_count += 1
if call_count < 3:
raise ConnectionError("DB connection refused")
return "ok"
result = await retry_db(flaky_op, max_retries=3, base_delay=0.01)
assert result == "ok"
assert call_count == 3
@pytest.mark.asyncio
async def test_non_transient_error_not_retried(self):
call_count = 0
async def failing_op():
nonlocal call_count
call_count += 1
raise ValueError("Invalid query")
with pytest.raises(ValueError):
await retry_db(failing_op, max_retries=3, base_delay=0.01)
assert call_count == 1
@pytest.mark.asyncio
async def test_all_retries_exhausted_raises_503(self):
from fastapi import HTTPException
async def always_fail():
raise ConnectionError("DB down")
with pytest.raises(HTTPException) as exc_info:
await retry_db(always_fail, max_retries=2, base_delay=0.01)
assert exc_info.value.status_code == 503
assert exc_info.value.detail["code"] == "db_unavailable"
@pytest.mark.asyncio
async def test_sqlalchemy_operational_error_retried(self):
from sqlalchemy.exc import OperationalError
call_count = 0
async def flaky_op():
nonlocal call_count
call_count += 1
if call_count < 2:
raise OperationalError("stmt", params={}, orig=Exception("connection reset"))
return "ok"
result = await retry_db(flaky_op, max_retries=3, base_delay=0.01)
assert result == "ok"
assert call_count == 2
class TestRedisFallback:
"""Tests for redis_call_with_fallback."""
@pytest.mark.asyncio
async def test_redis_success(self):
reset_all_circuits()
async def redis_get(key):
return f"value:{key}"
result = await redis_call_with_fallback(redis_get, None, "testkey")
assert result == "value:testkey"
@pytest.mark.asyncio
async def test_redis_failure_with_fallback(self):
reset_all_circuits()
async def redis_failing(key):
raise ConnectionError("Redis down")
async def db_fallback(key):
return f"db_value:{key}"
result = await redis_call_with_fallback(redis_failing, db_fallback, "testkey")
assert result == "db_value:testkey"
@pytest.mark.asyncio
async def test_redis_failure_no_fallback_returns_none(self):
reset_all_circuits()
async def redis_failing(key):
raise ConnectionError("Redis down")
result = await redis_call_with_fallback(redis_failing, None, "testkey")
assert result is None
@pytest.mark.asyncio
async def test_circuit_open_uses_fallback(self):
reset_all_circuits()
cb = get_circuit("redis")
# Force circuit open
cb._state = CircuitBreaker.OPEN
cb._opened_at = time.monotonic()
async def redis_op(key):
raise AssertionError("Should not be called when circuit is open")
async def fallback(key):
return f"fallback:{key}"
result = await redis_call_with_fallback(redis_op, fallback, "testkey")
assert result == "fallback:testkey"
class TestInMemoryRateLimiter:
"""Tests for InMemoryRateLimiter."""
@pytest.mark.asyncio
async def test_allows_under_limit(self):
limiter = InMemoryRateLimiter()
allowed, retry_after = await limiter.check("key1", max_attempts=5, window_seconds=60)
assert allowed is True
assert retry_after == 0
@pytest.mark.asyncio
async def test_blocks_over_limit(self):
limiter = InMemoryRateLimiter()
for _ in range(5):
await limiter.check("key1", max_attempts=5, window_seconds=60)
allowed, retry_after = await limiter.check("key1", max_attempts=5, window_seconds=60)
assert allowed is False
assert retry_after > 0
@pytest.mark.asyncio
async def test_different_keys_independent(self):
limiter = InMemoryRateLimiter()
for _ in range(3):
await limiter.check("key1", max_attempts=3, window_seconds=60)
allowed2, _ = await limiter.check("key2", max_attempts=3, window_seconds=60)
assert allowed2 is True
@pytest.mark.asyncio
async def test_reset_clears_key(self):
limiter = InMemoryRateLimiter()
for _ in range(3):
await limiter.check("key1", max_attempts=3, window_seconds=60)
await limiter.reset("key1")
allowed, _ = await limiter.check("key1", max_attempts=3, window_seconds=60)
assert allowed is True
@pytest.mark.asyncio
async def test_window_expiry(self):
limiter = InMemoryRateLimiter()
for _ in range(3):
await limiter.check("key1", max_attempts=3, window_seconds=1)
# Wait for window to expire
await asyncio.sleep(1.1)
allowed, _ = await limiter.check("key1", max_attempts=3, window_seconds=1)
assert allowed is True
def test_clear_all(self):
limiter = InMemoryRateLimiter()
limiter._counts["a"] = [1.0]
limiter._counts["b"] = [2.0]
limiter.clear_all()
assert len(limiter._counts) == 0
class TestCircuitBreakerMiddleware:
"""Tests for CircuitBreakerMiddleware ASGI behavior."""
@pytest.mark.asyncio
async def test_passes_through_when_closed(self):
reset_all_circuits()
called = False
async def app(scope, receive, send):
nonlocal called
called = True
await send({"type": "http.response.start", "status": 200, "headers": []})
await send({"type": "http.response.body", "body": b"ok"})
send_mock = AsyncMock()
middleware = CircuitBreakerMiddleware(app)
scope = {"type": "http", "path": "/api/v1/contacts"}
await middleware(scope, MagicMock(), send_mock)
assert called is True
@pytest.mark.asyncio
async def test_returns_503_when_open(self):
reset_all_circuits()
cb = get_circuit("db")
cb._state = CircuitBreaker.OPEN
cb._opened_at = time.monotonic()
called = False
async def app(scope, receive, send):
nonlocal called
called = True
sent_responses = []
async def mock_send(msg):
sent_responses.append(msg)
middleware = CircuitBreakerMiddleware(app)
scope = {"type": "http", "path": "/api/v1/contacts"}
await middleware(scope, MagicMock(), mock_send)
assert called is False
assert sent_responses[0]["status"] == 503
@pytest.mark.asyncio
async def test_skips_health_endpoints(self):
reset_all_circuits()
cb = get_circuit("db")
cb._state = CircuitBreaker.OPEN
cb._opened_at = time.monotonic()
called = False
async def app(scope, receive, send):
nonlocal called
called = True
await send({"type": "http.response.start", "status": 200, "headers": []})
await send({"type": "http.response.body", "body": b"ok"})
send_mock = AsyncMock()
middleware = CircuitBreakerMiddleware(app)
scope = {"type": "http", "path": "/api/v1/health"}
await middleware(scope, MagicMock(), send_mock)
assert called is True
@pytest.mark.asyncio
async def test_passes_websocket(self):
reset_all_circuits()
called = False
async def app(scope, receive, send):
nonlocal called
called = True
middleware = CircuitBreakerMiddleware(app)
scope = {"type": "websocket", "path": "/ws"}
await middleware(scope, MagicMock(), MagicMock())
assert called is True
# Fixtures: reset circuits before each test class
@pytest.fixture(autouse=True)
def reset_circuits():
reset_all_circuits()
get_inmemory_limiter().clear_all()
yield
reset_all_circuits()
get_inmemory_limiter().clear_all()