ForcePilot/backend/test/unit/channels/test_circuit_breaker.py

173 lines
5.0 KiB
Python
Raw Normal View History

from __future__ import annotations
import asyncio
import pytest
from yuxi.channels.infra.circuit_breaker import CircuitBreaker, CircuitBreakerOpenError, CircuitState
class TestCircuitBreakerInit:
def test_default_values(self):
cb = CircuitBreaker()
assert cb.failure_threshold == 5
assert cb.recovery_timeout == 60.0
assert cb.half_open_max_calls == 3
assert cb.state == CircuitState.CLOSED
assert cb._failure_count == 0
def test_custom_values(self):
cb = CircuitBreaker(failure_threshold=3, recovery_timeout=10.0, half_open_max_calls=2)
assert cb.failure_threshold == 3
assert cb.recovery_timeout == 10.0
assert cb.half_open_max_calls == 2
class TestCircuitStateTransitions:
@pytest.mark.asyncio
async def test_closed_to_open_on_threshold(self):
cb = CircuitBreaker(failure_threshold=3)
for _ in range(3):
await cb.record_failure()
assert cb.state == CircuitState.OPEN
@pytest.mark.asyncio
async def test_not_open_before_threshold(self):
cb = CircuitBreaker(failure_threshold=5)
for _ in range(3):
await cb.record_failure()
assert cb.state == CircuitState.CLOSED
@pytest.mark.asyncio
async def test_half_open_on_recovery_timeout(self):
cb = CircuitBreaker(failure_threshold=1, recovery_timeout=0.0)
await cb.record_failure()
assert cb.state == CircuitState.OPEN
async def dummy():
return "ok"
result = await cb.call(dummy)
assert result == "ok"
assert cb.state == CircuitState.CLOSED
@pytest.mark.asyncio
async def test_half_open_to_closed_on_success(self):
cb = CircuitBreaker(failure_threshold=1, recovery_timeout=0.0)
await cb.record_failure()
async def ok_func():
return "ok"
await cb.call(ok_func)
assert cb.state == CircuitState.CLOSED
@pytest.mark.asyncio
async def test_half_open_to_open_on_failure(self):
cb = CircuitBreaker(failure_threshold=1, recovery_timeout=0.0)
await cb.record_failure()
async def failing():
raise RuntimeError("boom")
with pytest.raises(RuntimeError, match="boom"):
await cb.call(failing)
assert cb.state == CircuitState.OPEN
@pytest.mark.asyncio
async def test_record_success_resets_closed_count(self):
cb = CircuitBreaker(failure_threshold=5)
for _ in range(3):
await cb.record_failure()
await cb.record_success()
assert cb._failure_count == 0
assert cb.state == CircuitState.CLOSED
class TestCircuitBreakerCall:
@pytest.mark.asyncio
async def test_call_passes_result(self):
cb = CircuitBreaker()
async def func():
return "hello"
result = await cb.call(func)
assert result == "hello"
@pytest.mark.asyncio
async def test_call_records_failure_on_exception(self):
cb = CircuitBreaker()
async def func():
raise RuntimeError("fail")
with pytest.raises(RuntimeError, match="fail"):
await cb.call(func)
assert cb._failure_count == 1
@pytest.mark.asyncio
async def test_call_blocks_when_open(self):
cb = CircuitBreaker(failure_threshold=1, recovery_timeout=999.0)
await cb.record_failure()
assert cb.state == CircuitState.OPEN
async def ok_func():
return "ok"
with pytest.raises(CircuitBreakerOpenError):
await cb.call(ok_func)
@pytest.mark.asyncio
async def test_call_recovers_after_timeout(self):
cb = CircuitBreaker(failure_threshold=1, recovery_timeout=0.0)
await cb.record_failure()
async def ok_func():
return "ok"
await cb.call(ok_func)
assert cb.state == CircuitState.CLOSED
class TestHalfOpenMaxCalls:
@pytest.mark.asyncio
async def test_half_open_max_calls_reached(self):
cb = CircuitBreaker(half_open_max_calls=2)
cb.state = CircuitState.HALF_OPEN
cb._half_open_calls = 2
async def ok_func():
return "ok"
with pytest.raises(CircuitBreakerOpenError):
await cb.call(ok_func)
class TestConcurrentSafety:
@pytest.mark.asyncio
async def test_concurrent_failures(self):
cb = CircuitBreaker(failure_threshold=3)
async with cb._lock:
pass
results = await asyncio.gather(
cb.record_failure(),
cb.record_failure(),
cb.record_failure(),
)
assert cb._failure_count == 3
@pytest.mark.asyncio
async def test_concurrent_record_success(self):
cb = CircuitBreaker(failure_threshold=5)
for _ in range(3):
await cb.record_failure()
assert cb._failure_count == 3
await asyncio.gather(
cb.record_success(),
cb.record_success(),
cb.record_success(),
)
assert cb._failure_count == 0
assert cb.state == CircuitState.CLOSED