Canary Deployments for ML Models: Shadow Testing and Progressive Rollouts
Deploying a new ML model version is riskier than deploying a new API version. A bad API deployment returns errors — observable, alertable, rollback-triggering. A bad model deployment returns predictions that look valid but are wrong. Accuracy can degrade 20% before a business metric moves enough to alert. By then you've made thousands of bad decisions.
Canary deployments with shadow testing solve this by letting you validate a new model against real traffic before it makes any production decisions.
The Core Pattern
Shadow mode and canary deployment are different things:
Shadow mode: New model receives a copy of all production traffic and makes predictions, but those predictions are discarded. Zero risk. Used to measure accuracy before any traffic switch.
Canary deployment: A small percentage of production traffic (1–5%) routes to the new model and those predictions are actually served. Used to measure business outcomes at low risk before full rollout.
Run shadow mode first, then canary, then full rollout.
Implementing Shadow Mode
The shadow proxy intercepts requests, forwards them to both models, logs both responses, and returns only the production model's response.
# shadow_proxy.py
import asyncio
import aiohttp
import logging
import time
from dataclasses import dataclass, asdict
from typing import Optional
import json
import uuid
logger = logging.getLogger(__name__)
@dataclass
class ShadowComparison:
request_id: str
timestamp: float
input_hash: str
production_prediction: list
shadow_prediction: list
production_latency_ms: float
shadow_latency_ms: float
agreement: bool
agreement_threshold: float
async def shadow_predict(
session: aiohttp.ClientSession,
production_url: str,
shadow_url: str,
payload: dict,
agreement_threshold: float = 0.95,
) -> tuple[dict, ShadowComparison]:
"""
Send request to both models. Return production response.
Log comparison for analysis.
"""
import hashlib
input_hash = hashlib.md5(json.dumps(payload, sort_keys=True).encode()).hexdigest()
request_id = str(uuid.uuid4())
async def fetch(url: str) -> tuple[dict, float]:
start = time.perf_counter()
async with session.post(url, json=payload) as resp:
result = await resp.json()
latency_ms = (time.perf_counter() - start) * 1000
return result, latency_ms
# Run both concurrently
(prod_result, prod_latency), (shadow_result, shadow_latency) = await asyncio.gather(
fetch(production_url),
fetch(shadow_url),
)
prod_preds = prod_result.get("predictions", [])
shadow_preds = shadow_result.get("predictions", [])
# Agreement check: are top-1 predictions the same?
import numpy as np
prod_labels = [np.argmax(p) for p in prod_preds]
shadow_labels = [np.argmax(p) for p in shadow_preds]
n = len(prod_labels)
if n > 0:
agreement_rate = sum(p == s for p, s in zip(prod_labels, shadow_labels)) / n
else:
agreement_rate = 1.0
comparison = ShadowComparison(
request_id=request_id,
timestamp=time.time(),
input_hash=input_hash,
production_prediction=prod_labels,
shadow_prediction=shadow_labels,
production_latency_ms=prod_latency,
shadow_latency_ms=shadow_latency,
agreement=agreement_rate >= agreement_threshold,
agreement_threshold=agreement_threshold,
)
# Log asynchronously — don't block the response
logger.info("shadow_comparison", extra=asdict(comparison))
return prod_result, comparisonTest the shadow proxy:
# tests/test_shadow_proxy.py
import pytest
import asyncio
import aiohttp
from unittest.mock import AsyncMock, patch, MagicMock
import json
import numpy as np
PROD_URL = "http://localhost:8501/v1/models/model_v1:predict"
SHADOW_URL = "http://localhost:8502/v1/models/model_v2:predict"
def make_mock_response(predictions: list) -> AsyncMock:
mock_resp = AsyncMock()
mock_resp.json = AsyncMock(return_value={"predictions": predictions})
mock_resp.__aenter__ = AsyncMock(return_value=mock_resp)
mock_resp.__aexit__ = AsyncMock(return_value=None)
return mock_resp
class TestShadowProxy:
@pytest.mark.asyncio
async def test_returns_production_response_not_shadow(self):
prod_pred = [[0.9, 0.1]]
shadow_pred = [[0.1, 0.9]]
with patch("aiohttp.ClientSession.post") as mock_post:
mock_post.side_effect = [
make_mock_response(prod_pred),
make_mock_response(shadow_pred),
]
async with aiohttp.ClientSession() as session:
result, comparison = await shadow_predict(
session, PROD_URL, SHADOW_URL,
payload={"instances": [[1.0, 2.0, 3.0, 4.0]]}
)
assert result["predictions"] == prod_pred
@pytest.mark.asyncio
async def test_agreement_when_predictions_match(self):
same_pred = [[0.9, 0.1]]
with patch("aiohttp.ClientSession.post") as mock_post:
mock_post.side_effect = [
make_mock_response(same_pred),
make_mock_response(same_pred),
]
async with aiohttp.ClientSession() as session:
_, comparison = await shadow_predict(
session, PROD_URL, SHADOW_URL,
payload={"instances": [[1.0, 2.0, 3.0, 4.0]]}
)
assert comparison.agreement is True
@pytest.mark.asyncio
async def test_disagreement_when_predictions_differ(self):
prod_pred = [[0.9, 0.1]] # label 0
shadow_pred = [[0.1, 0.9]] # label 1
with patch("aiohttp.ClientSession.post") as mock_post:
mock_post.side_effect = [
make_mock_response(prod_pred),
make_mock_response(shadow_pred),
]
async with aiohttp.ClientSession() as session:
_, comparison = await shadow_predict(
session, PROD_URL, SHADOW_URL,
payload={"instances": [[1.0, 2.0, 3.0, 4.0]]}
)
assert comparison.agreement is False
@pytest.mark.asyncio
async def test_shadow_failure_does_not_affect_production_response(self):
"""If shadow model is down, production still returns normally."""
prod_pred = [[0.8, 0.2]]
async def mock_post(url, **kwargs):
if "8502" in url: # shadow URL
raise aiohttp.ClientConnectionError("Shadow model down")
return make_mock_response(prod_pred)
with patch("aiohttp.ClientSession.post", side_effect=mock_post):
async with aiohttp.ClientSession() as session:
try:
result, _ = await shadow_predict(
session, PROD_URL, SHADOW_URL,
payload={"instances": [[1.0, 2.0, 3.0, 4.0]]}
)
assert result["predictions"] == prod_pred
except aiohttp.ClientConnectionError:
pytest.fail("Shadow failure should not propagate to caller")Traffic Splitting for A/B Testing
Once shadow mode shows acceptable agreement, switch to canary with real traffic split:
# canary_router.py
import random
import hashlib
from typing import Callable, Any
from dataclasses import dataclass
from enum import Enum
class ModelVersion(Enum):
PRODUCTION = "v1"
CANARY = "v2"
@dataclass
class CanaryConfig:
canary_percentage: float # 0.0 to 1.0
sticky: bool = True # Same user always gets same model
def select_model_version(
user_id: str,
config: CanaryConfig,
) -> ModelVersion:
"""
Deterministic model selection based on user_id hash.
Same user always gets same model version (sticky routing).
"""
if not config.sticky:
return ModelVersion.CANARY if random.random() < config.canary_percentage else ModelVersion.PRODUCTION
# Hash-based: deterministic per user_id
hash_value = int(hashlib.md5(user_id.encode()).hexdigest(), 16)
bucket = (hash_value % 100) / 100.0 # 0.0 to 1.0
return ModelVersion.CANARY if bucket < config.canary_percentage else ModelVersion.PRODUCTION
def route_prediction_request(
user_id: str,
payload: dict,
config: CanaryConfig,
production_endpoint: Callable,
canary_endpoint: Callable,
) -> tuple[Any, ModelVersion]:
version = select_model_version(user_id, config)
if version == ModelVersion.CANARY:
result = canary_endpoint(payload)
else:
result = production_endpoint(payload)
return result, versionTest the router:
# tests/test_canary_router.py
import pytest
from collections import Counter
from canary_router import select_model_version, ModelVersion, CanaryConfig
class TestCanaryRouter:
def test_5_percent_canary_routes_approximately_5_percent(self):
config = CanaryConfig(canary_percentage=0.05)
n_samples = 10_000
user_ids = [f"user_{i}" for i in range(n_samples)]
versions = [select_model_version(uid, config) for uid in user_ids]
canary_count = sum(v == ModelVersion.CANARY for v in versions)
canary_rate = canary_count / n_samples
# Allow ±2% tolerance
assert 0.03 <= canary_rate <= 0.07, (
f"Expected ~5% canary traffic, got {canary_rate:.1%}"
)
def test_sticky_routing_same_user_same_version(self):
config = CanaryConfig(canary_percentage=0.1, sticky=True)
user_id = "user_42"
versions = [select_model_version(user_id, config) for _ in range(100)]
assert len(set(versions)) == 1, (
"Sticky routing: same user should always get same model version"
)
def test_non_sticky_routing_can_vary(self):
config = CanaryConfig(canary_percentage=0.5, sticky=False)
user_id = "user_42"
versions = [select_model_version(user_id, config) for _ in range(1000)]
unique_versions = set(versions)
# With 50% probability and 1000 samples, both versions should appear
assert len(unique_versions) == 2, (
"Non-sticky routing should sometimes route to canary"
)
def test_0_percent_canary_always_routes_to_production(self):
config = CanaryConfig(canary_percentage=0.0)
for i in range(100):
version = select_model_version(f"user_{i}", config)
assert version == ModelVersion.PRODUCTION
def test_100_percent_canary_always_routes_to_canary(self):
config = CanaryConfig(canary_percentage=1.0)
for i in range(100):
version = select_model_version(f"user_{i}", config)
assert version == ModelVersion.CANARYRollback Strategy Testing
A canary deployment without a tested rollback is a liability. Test the rollback mechanism:
# canary_controller.py
from dataclasses import dataclass, field
from typing import Optional
from datetime import datetime
import threading
@dataclass
class CanaryMetrics:
canary_error_rate: float = 0.0
production_error_rate: float = 0.0
canary_p99_latency_ms: float = 0.0
production_p99_latency_ms: float = 0.0
canary_accuracy: Optional[float] = None
production_accuracy: Optional[float] = None
@dataclass
class RollbackPolicy:
max_error_rate_delta: float = 0.02 # Canary error rate can be at most 2% higher
max_latency_delta_ms: float = 50.0 # Canary p99 can be at most 50ms slower
min_accuracy_delta: float = -0.02 # Canary accuracy can drop at most 2%
class CanaryController:
def __init__(self, policy: RollbackPolicy):
self.policy = policy
self._canary_percentage = 0.0
self._rolled_back = False
self._lock = threading.Lock()
def should_rollback(self, metrics: CanaryMetrics) -> tuple[bool, str]:
"""Returns (should_rollback, reason)."""
error_delta = metrics.canary_error_rate - metrics.production_error_rate
if error_delta > self.policy.max_error_rate_delta:
return True, (
f"Canary error rate {metrics.canary_error_rate:.1%} exceeds "
f"production by {error_delta:.1%} (threshold: {self.policy.max_error_rate_delta:.1%})"
)
latency_delta = metrics.canary_p99_latency_ms - metrics.production_p99_latency_ms
if latency_delta > self.policy.max_latency_delta_ms:
return True, (
f"Canary p99 latency {metrics.canary_p99_latency_ms:.0f}ms exceeds "
f"production by {latency_delta:.0f}ms (threshold: {self.policy.max_latency_delta_ms:.0f}ms)"
)
if (metrics.canary_accuracy is not None and metrics.production_accuracy is not None):
accuracy_delta = metrics.canary_accuracy - metrics.production_accuracy
if accuracy_delta < self.policy.min_accuracy_delta:
return True, (
f"Canary accuracy {metrics.canary_accuracy:.3f} dropped "
f"{abs(accuracy_delta):.3f} below production (threshold: {abs(self.policy.min_accuracy_delta):.3f})"
)
return False, ""
def rollback(self) -> None:
with self._lock:
self._canary_percentage = 0.0
self._rolled_back = True
@property
def is_rolled_back(self) -> bool:
return self._rolled_back
# tests/test_rollback.py
class TestRollbackPolicy:
@pytest.fixture
def controller(self):
return CanaryController(policy=RollbackPolicy(
max_error_rate_delta=0.02,
max_latency_delta_ms=50.0,
min_accuracy_delta=-0.02,
))
def test_healthy_metrics_no_rollback(self, controller):
metrics = CanaryMetrics(
canary_error_rate=0.01,
production_error_rate=0.01,
canary_p99_latency_ms=80.0,
production_p99_latency_ms=75.0,
canary_accuracy=0.92,
production_accuracy=0.91,
)
should_rb, reason = controller.should_rollback(metrics)
assert not should_rb
def test_high_error_rate_triggers_rollback(self, controller):
metrics = CanaryMetrics(
canary_error_rate=0.08, # 8%
production_error_rate=0.01, # 1%
canary_p99_latency_ms=75.0,
production_p99_latency_ms=75.0,
)
should_rb, reason = controller.should_rollback(metrics)
assert should_rb
assert "error rate" in reason.lower()
def test_high_latency_triggers_rollback(self, controller):
metrics = CanaryMetrics(
canary_error_rate=0.01,
production_error_rate=0.01,
canary_p99_latency_ms=200.0, # 125ms above production
production_p99_latency_ms=75.0,
)
should_rb, reason = controller.should_rollback(metrics)
assert should_rb
assert "latency" in reason.lower()
def test_accuracy_drop_triggers_rollback(self, controller):
metrics = CanaryMetrics(
canary_error_rate=0.01,
production_error_rate=0.01,
canary_p99_latency_ms=75.0,
production_p99_latency_ms=75.0,
canary_accuracy=0.85, # 6% drop
production_accuracy=0.91,
)
should_rb, reason = controller.should_rollback(metrics)
assert should_rb
assert "accuracy" in reason.lower()
def test_rollback_sets_canary_percentage_to_zero(self, controller):
controller._canary_percentage = 0.10
controller.rollback()
assert controller._canary_percentage == 0.0
assert controller.is_rolled_backProgressive Rollout Pattern
Test the rollout schedule logic:
# rollout_scheduler.py
from typing import List
from dataclasses import dataclass
@dataclass
class RolloutStage:
percentage: float
min_duration_minutes: int
required_sample_size: int
STANDARD_ROLLOUT = [
RolloutStage(percentage=0.01, min_duration_minutes=30, required_sample_size=1000),
RolloutStage(percentage=0.05, min_duration_minutes=60, required_sample_size=5000),
RolloutStage(percentage=0.10, min_duration_minutes=120, required_sample_size=10000),
RolloutStage(percentage=0.25, min_duration_minutes=240, required_sample_size=25000),
RolloutStage(percentage=1.00, min_duration_minutes=0, required_sample_size=0),
]
def can_advance_stage(
current_stage: RolloutStage,
elapsed_minutes: float,
samples_collected: int,
metrics: CanaryMetrics,
policy: RollbackPolicy,
) -> tuple[bool, str]:
if elapsed_minutes < current_stage.min_duration_minutes:
return False, (
f"Need {current_stage.min_duration_minutes - elapsed_minutes:.0f} more minutes"
)
if samples_collected < current_stage.required_sample_size:
return False, (
f"Need {current_stage.required_sample_size - samples_collected} more samples"
)
controller = CanaryController(policy)
should_rb, reason = controller.should_rollback(metrics)
if should_rb:
return False, f"Metrics failing: {reason}"
return True, "Ready to advance"
# tests/test_rollout_scheduler.py
class TestProgressiveRollout:
GOOD_METRICS = CanaryMetrics(
canary_error_rate=0.01, production_error_rate=0.01,
canary_p99_latency_ms=75.0, production_p99_latency_ms=75.0,
)
POLICY = RollbackPolicy()
def test_cannot_advance_before_min_duration(self):
stage = RolloutStage(percentage=0.01, min_duration_minutes=30, required_sample_size=100)
can_advance, reason = can_advance_stage(stage, 15.0, 5000, self.GOOD_METRICS, self.POLICY)
assert not can_advance
assert "15 more minutes" in reason
def test_cannot_advance_without_enough_samples(self):
stage = RolloutStage(percentage=0.01, min_duration_minutes=30, required_sample_size=1000)
can_advance, reason = can_advance_stage(stage, 60.0, 100, self.GOOD_METRICS, self.POLICY)
assert not can_advance
assert "more samples" in reason
def test_can_advance_when_all_criteria_met(self):
stage = RolloutStage(percentage=0.01, min_duration_minutes=30, required_sample_size=1000)
can_advance, reason = can_advance_stage(stage, 60.0, 2000, self.GOOD_METRICS, self.POLICY)
assert can_advance
def test_cannot_advance_when_metrics_failing(self):
bad_metrics = CanaryMetrics(
canary_error_rate=0.10, production_error_rate=0.01,
canary_p99_latency_ms=75.0, production_p99_latency_ms=75.0,
)
stage = RolloutStage(percentage=0.01, min_duration_minutes=30, required_sample_size=1000)
can_advance, reason = can_advance_stage(stage, 60.0, 2000, bad_metrics, self.POLICY)
assert not can_advance
assert "Metrics failing" in reasonBusiness Metric Monitoring
Technical metrics (error rate, latency) aren't enough. Test that business metric monitoring is wired correctly:
# business_metrics.py
from dataclasses import dataclass
from typing import Optional
@dataclass
class BusinessMetrics:
"""Business outcomes for a given time window."""
conversion_rate: float
revenue_per_session: float
click_through_rate: float
sample_size: int
def compute_business_impact(
canary: BusinessMetrics,
production: BusinessMetrics,
min_sample_size: int = 1000,
) -> dict:
"""Compute relative lift/drop and statistical significance."""
from scipy import stats
import math
if canary.sample_size < min_sample_size or production.sample_size < min_sample_size:
return {"significant": False, "reason": "Insufficient sample size"}
# Z-test for conversion rate
p1 = canary.conversion_rate
p2 = production.conversion_rate
n1 = canary.sample_size
n2 = production.sample_size
p_pool = (p1 * n1 + p2 * n2) / (n1 + n2)
se = math.sqrt(p_pool * (1 - p_pool) * (1/n1 + 1/n2))
z_stat = (p1 - p2) / se if se > 0 else 0
p_value = 2 * (1 - stats.norm.cdf(abs(z_stat)))
return {
"conversion_lift": (p1 - p2) / p2 if p2 > 0 else 0,
"z_statistic": z_stat,
"p_value": p_value,
"significant": p_value < 0.05,
"direction": "positive" if p1 > p2 else "negative",
}Summary
A production-ready canary deployment for ML models requires:
- Shadow mode first — run the new model on real traffic, compare predictions to production, measure disagreement rate. Only proceed if disagreement is within tolerance.
- Sticky routing — hash-based traffic splitting ensures the same user always gets the same model, preventing inconsistent user experiences.
- Automatic rollback — define rollback triggers for error rate, latency, and accuracy. Test both the trigger conditions and the rollback mechanism itself.
- Progressive stages — enforce minimum duration and sample size at each stage. Don't advance on time alone.
- Business metrics — conversion rate and revenue per session catch what technical metrics miss. Test the statistical significance logic so you don't act on noise.
For end-to-end testing of canary deployment workflows — including verifying that shadow mode comparison logic, rollback triggers, and progressive rollout gates work correctly across staging environments — HelpMeTest provides the orchestration and monitoring layer without custom infrastructure.