Topic 7: Scale and Reliability
16 min read·21 Sept 2026
module12/reliability.py
python
# module12/reliability.py
"""Staying useful when parts break: timeouts, circuit breakers, fallback ladders, load tests."""
from __future__ import annotations
import random
import statistics
import time
from concurrent.futures import ThreadPoolExecutor, TimeoutError as FutureTimeout
from dataclasses import dataclass, field
from enum import Enum
class Health(str, Enum):
OK = "ok"
DEGRADED = "degraded"
DOWN = "down"
class CircuitOpen(Exception):
"""The dependency is failing; stop calling it for a while."""
@dataclass
class CircuitBreaker:
"""After N failures, stop calling a broken dependency until a cool-off passes."""
name: str
failure_threshold: int = 3
cool_off_seconds: float = 10.0
failures: int = 0
opened_at: float | None = None
@property
def state(self) -> str:
if self.opened_at is None:
return "closed"
return "half-open" if time.monotonic() - self.opened_at > self.cool_off_seconds else "open"
def call(self, fn, *args, **kwargs):
if self.state == "open":
raise CircuitOpen(f"{self.name} circuit open ({self.failures} failures)")
try:
result = fn(*args, **kwargs)
except Exception:
self.failures += 1
if self.failures >= self.failure_threshold:
self.opened_at = time.monotonic()
raise
self.failures, self.opened_at = 0, None # success closes the circuit again
return result
def with_timeout(fn, seconds: float, *args, **kwargs):
"""Run a call with a hard deadline. A slow dependency must not hold the whole request."""
with ThreadPoolExecutor(max_workers=1) as pool:
future = pool.submit(fn, *args, **kwargs)
try:
return future.result(timeout=seconds)
except FutureTimeout:
raise TimeoutError(f"call exceeded {seconds}s")
@dataclass
class FallbackLadder:
"""Try the best option; on failure, step down. Record which rung answered."""
name: str
rungs: list[tuple[str, object]] = field(default_factory=list) # (label, callable)
def run(self, *args, **kwargs) -> tuple[object, str, list[str]]:
problems = []
for label, step in self.rungs:
try:
return step(*args, **kwargs), label, problems
except Exception as err:
problems.append(f"{label}: {type(err).__name__}: {err}")
raise RuntimeError(f"{self.name}: every fallback failed: {problems}")
# ---------------------------------------------------------------- degradation policy
DEGRADATION = {
"reranker": ("skip reranking, send the first-stage top k",
"slightly worse ranking; answer still grounded"),
"vector_store": ("fall back to BM25 (or a cached index snapshot)",
"keyword-only recall; paraphrases may miss"),
"embedding_api": ("serve cached query embeddings, else BM25 only",
"new phrasings degrade to keyword search"),
"llm_provider": ("switch provider, then a smaller model, then extractive answer",
"an extractive answer quotes passages without composing them"),
"query_transformer": ("use the raw query", "worse recall on messy queries"),
"database": ("answer document questions only, say order data is unavailable",
"partial answer with an honest note"),
}
def extractive_answer(question: str, passages: list[tuple[str, str]], limit: int = 2) -> str:
"""Last rung of the ladder: no LLM, so quote the best passages and say so."""
if not passages:
return ("Our assistant is temporarily degraded and I could not find a relevant document. "
"Please try again shortly or contact support.")
quoted = " ".join(f"[{doc_id}] {text.strip()}" for doc_id, text in passages[:limit])
return ("Our assistant is running in a reduced mode, so here are the most relevant passages "
f"instead of a written answer: {quoted}")
# ---------------------------------------------------------------- load testing
@dataclass
class LoadTestResult:
requests: int
errors: int
seconds: float
latencies_ms: list[float]
def summary(self) -> dict:
ordered = sorted(self.latencies_ms)
def pct(p):
return round(ordered[min(len(ordered) - 1, int(len(ordered) * p / 100))], 1) if ordered else 0.0
return {"requests": self.requests, "errors": self.errors,
"error_rate": round(self.errors / self.requests, 3) if self.requests else 0.0,
"throughput_rps": round(self.requests / self.seconds, 1) if self.seconds else 0.0,
"p50_ms": pct(50), "p95_ms": pct(95), "p99_ms": pct(99),
"mean_ms": round(statistics.mean(ordered), 1) if ordered else 0.0}
def load_test(handler, queries: list[str], concurrency: int = 8, total: int = 100,
seed: int = 0) -> LoadTestResult:
"""Send `total` requests with `concurrency` workers and report the percentiles users feel."""
rng = random.Random(seed)
work = [rng.choice(queries) for _ in range(total)]
latencies, errors = [], 0
start = time.perf_counter()
def run_one(query: str):
began = time.perf_counter()
try:
handler(query)
return (time.perf_counter() - began) * 1000, None
except Exception as err:
return (time.perf_counter() - began) * 1000, err
with ThreadPoolExecutor(max_workers=concurrency) as pool:
for latency, error in pool.map(run_one, work):
latencies.append(latency)
errors += bool(error)
return LoadTestResult(total, errors, time.perf_counter() - start, latencies)
class FlakyDependency:
"""A stand-in for a real dependency, so you can test degradation without breaking production."""
def __init__(self, name: str, fail_rate: float = 0.0, latency_ms: float = 0.0,
down: bool = False, seed: int = 0):
self.name, self.fail_rate, self.latency_ms, self.down = name, fail_rate, latency_ms, down
self.rng = random.Random(seed)
self.calls = 0
def __call__(self, *args, **kwargs):
self.calls += 1
if self.latency_ms:
time.sleep(self.latency_ms / 1000)
if self.down or self.rng.random() < self.fail_rate:
raise ConnectionError(f"{self.name} unavailable")
return f"{self.name} ok"