Course Content
Model Deployment for AI Engineers
4 sections · 10 lessons
Security and Scalability — Hardening a Model API
A credit-scoring API goes live with no authentication. It is on a private subnet, so the team reasons that only internal services can reach it. Eight weeks later, a competitor launches a scoring product with suspiciously similar behaviour on edge cases.
The access logs tell the story. One IP address made 340,000 requests over six weeks — about one every eleven seconds, a trickle low enough that no per-minute alert ever fired. Each request was a synthetic applicant profile sweeping systematically across income, age, and debt ratio. With 340,000 input-output pairs the attacker had enough labelled data to train a surrogate model that reproduced the original's decisions with 94% agreement.
Nothing was breached. No credentials were stolen, no database was accessed, no vulnerability was exploited. The API did exactly what it was built to do, 340,000 times, and in doing so it gave away the model. This is the thing that makes securing an ML service different: the intended functionality is itself the attack surface.
The threat model
A model-serving API inherits every web-API threat and adds several of its own.
| Threat | What the attacker does | Primary control |
|---|---|---|
| Unauthorised access | Calls the endpoint directly | Authentication on every route |
| Model extraction | Systematic queries to clone the model | Rate limits per identity, query-pattern detection, coarse outputs |
| Membership inference | Determines whether a specific record was in your training data | Avoid returning raw confidences; differential privacy in training |
| Adversarial input | Crafts an input that forces a chosen wrong output | Input validation, ensembling, adversarial training |
| Resource exhaustion | Huge payloads or unbounded batches | Size caps, batch caps, timeouts, quotas |
| Data exfiltration via logs | Reads personal data your service logged | Log hashes and shapes, never raw inputs |
| Prompt injection (generative models) | Embeds instructions in user text that the model obeys | Separate instruction and data channels; validate outputs |
Two of these deserve arithmetic, because the numbers change how seriously you take them.
Model extraction is cheap. Published attacks recover a usable copy of a moderately sized classifier in roughly 105 queries. At an unthrottled 10 requests per second, 100,000 queries take 100000 / 10 / 3600 = 2.8 hours. With a per-key quota of 1,000 requests per day, the same attack takes 100 days and leaves an obvious trail. Rate limiting is not just a capacity control; against extraction it is the primary defence.
Returning full probability vectors leaks more than you think. A label alone gives the attacker one bit per query. A full softmax over four classes gives a real-valued gradient signal, which is why extraction attacks using confidences need one to two orders of magnitude fewer queries than attacks using labels alone. If your product does not need the vector, do not return it — round to two decimal places, or return the top label with a coarse bucket.
Every extra digit of confidence you return is free training signal for whoever wants to clone your model.
Authentication: who is allowed to call this
API keys
Right for service-to-service calls. The critical detail is storage: keys are credentials, so hash them exactly as you would passwords.
1import hashlib, hmac, secrets2from fastapi import Security, HTTPException, status3from fastapi.security import APIKeyHeader45api_key_header = APIKeyHeader(name="X-API-Key", auto_error=False)67def issue_key() -> tuple[str, str]:8 raw = "sk_" + secrets.token_urlsafe(32) # shown to the user exactly once9 return raw, hashlib.sha256(raw.encode()).hexdigest()1011async def require_key(key: str = Security(api_key_header)):12 if not key:13 raise HTTPException(status.HTTP_401_UNAUTHORIZED, "missing X-API-Key")14 digest = hashlib.sha256(key.encode()).hexdigest()15 record = await db.fetch_one(16 "SELECT * FROM api_keys WHERE key_hash = :h AND revoked_at IS NULL",17 {"h": digest})18 if record is None or not hmac.compare_digest(record["key_hash"], digest):19 raise HTTPException(status.HTTP_401_UNAUTHORIZED, "invalid key")20 if record["expires_at"] < utcnow():21 raise HTTPException(status.HTTP_401_UNAUTHORIZED, "key expired")22 return record2324@app.post("/predict")25async def predict(req: PredictRequest, client = Depends(require_key)):26 ...Three things here matter and are commonly skipped. Storing the hash rather than the key means a database leak does not hand over working credentials. hmac.compare_digest is a constant-time comparison — a plain == returns faster on an early mismatch, which over many attempts leaks the prefix. And an expiry plus a revocation column is what lets you turn off a leaked key in seconds rather than rotating everything.
JWT bearer tokens
Right when a user-facing application calls you and you need identity and permissions carried in the request. A JWT is a signed JSON payload; your service verifies the signature and reads the claims without a database round trip.
1import jwt2from datetime import datetime, timedelta, timezone3from fastapi.security import HTTPBearer, HTTPAuthorizationCredentials45SECRET, ALGO = os.environ["JWT_SECRET"], "HS256"6bearer = HTTPBearer()78def create_token(user_id: str, scopes: list[str], minutes: int = 30) -> str:9 now = datetime.now(timezone.utc)10 return jwt.encode({11 "sub": user_id, "scopes": scopes,12 "iat": now, "exp": now + timedelta(minutes=minutes),13 "iss": "auth.acme.com", "aud": "inference-api",14 }, SECRET, algorithm=ALGO)1516def require_scope(scope: str): # plain def: called at import time17 async def dependency(cred: HTTPAuthorizationCredentials = Security(bearer)):18 try:19 claims = jwt.decode(cred.credentials, SECRET,20 algorithms=[ALGO], # never accept "none"21 audience="inference-api",22 issuer="auth.acme.com")23 except jwt.ExpiredSignatureError:24 raise HTTPException(401, "token expired")25 except jwt.InvalidTokenError as e:26 raise HTTPException(401, f"invalid token: {e}")27 if scope not in claims.get("scopes", []):28 raise HTTPException(403, f"missing scope: {scope}")29 return claims30 return dependency3132@app.post("/predict")33async def predict(req: PredictRequest, claims = Depends(require_scope("predict:write"))):34 ...Pinning algorithms=[ALGO] is not stylistic. The classic JWT vulnerability is a library that honours the alg header from the token itself: an attacker sets alg: none, strips the signature, and the library accepts it. Passing an explicit allow-list closes that. Checking audience and issuer closes the related attack where a valid token minted for a different service is replayed against yours.
The trade-off to understand: JWTs are stateless, so you cannot revoke one before it expires without adding a denylist — which reintroduces the database lookup you were avoiding. Keep expiry short (15–30 minutes) and use refresh tokens.
| API key | JWT | mTLS | |
|---|---|---|---|
| Caller | Another service | An end user via an app | Another service, high trust |
| Carries permissions | Via a database lookup | In the token itself | In the certificate |
| Revocation | Immediate | Only at expiry, or via a denylist | Via CRL / short-lived certs |
| Per-request cost | One DB or cache lookup | Signature verification only | TLS handshake, then free |
| Main risk | Key committed to a repository | Long expiry, unpinned algorithm | Certificate management burden |
Input validation is a security control
Schema validation is usually framed as data hygiene. In a model-serving API it is also a defence, because unbounded inputs are the cheapest denial-of-service available.
1from pydantic import BaseModel, Field, field_validator2from typing import Literal, List3import math45class Applicant(BaseModel):6 annual_income: float = Field(..., ge=0, le=10_000_000)7 age: int = Field(..., ge=18, le=120)8 debt_ratio: float = Field(..., ge=0, le=5.0)9 employment: Literal["full_time", "part_time", "self_employed", "retired"]1011 @field_validator("annual_income", "debt_ratio")12 @classmethod13 def finite(cls, v: float) -> float:14 if not math.isfinite(v): # NaN and inf pass float() happily15 raise ValueError("value must be finite")16 return v1718class BatchRequest(BaseModel):19 items: List[Applicant] = Field(..., min_length=1, max_length=50)The max_length=50 is the load-bearing line. Without it, one request containing 200,000 items allocates until the container is killed — a complete outage from a single well-formed request. The finiteness check matters too: JSON permits values that parse to NaN and inf in Python, and a NaN propagates silently through a model to produce a NaN score, which downstream comparisons treat as "not greater than the threshold" — quietly approving everything.
Rate limiting that holds across replicas
An in-process counter breaks the moment you run more than one replica: four replicas each allowing 100 per minute is a real limit of 400, and it resets on every deploy. Shared state is required.
1import redis.asyncio as redis23r = redis.from_url(os.environ["REDIS_URL"])45# Atomic sliding window. Executed server-side so concurrent calls cannot interleave.6SLIDING_WINDOW = """7local key, now, window, limit = KEYS[1], tonumber(ARGV[1]), tonumber(ARGV[2]), tonumber(ARGV[3])8redis.call('ZREMRANGEBYSCORE', key, 0, now - window)9local used = redis.call('ZCARD', key)10if used < limit then11 redis.call('ZADD', key, now, ARGV[4])12 redis.call('EXPIRE', key, window)13 return {1, limit - used - 1}14end15return {0, 0}16"""17_script = r.register_script(SLIDING_WINDOW)1819TIERS = {"free": (100, 3600), "pro": (10_000, 3600), "enterprise": (200_000, 3600)}2021async def enforce_quota(client = Depends(require_key)):22 limit, window = TIERS[client["tier"]]23 allowed, remaining = await _script(24 keys=[f"rl:{client['id']}"],25 args=[time.time(), window, limit, uuid.uuid4().hex])26 if not allowed:27 raise HTTPException(429, "rate limit exceeded",28 headers={"Retry-After": str(window)})29 return {"remaining": remaining}A sliding window is worth the extra complexity over a fixed window. With fixed hourly windows and a limit of 100, a client can send 100 at 10:59 and 100 at 11:01 — 200 requests in two minutes, double the intended rate, at exactly the boundary. The sliding window counts the trailing hour continuously, so that trick does not work.
Layer a second limit on top for extraction defence: a daily cap and an alert on query diversity. A legitimate client's requests cluster around a few real distributions. An extraction attack sweeps a grid, so its inputs are unusually uniform across feature space. Tracking the entropy of a client's inputs over a day flags that pattern well before the 340,000th request.
Transport, and what belongs in front of your app
Never terminate TLS in your Python process. A reverse proxy in front does TLS, buffering, static rate limiting, and request-size enforcement far more efficiently, and it stops slow-client attacks from consuming your worker threads.
upstream inference { least_conn; server api-1:8000 max_fails=3 fail_timeout=30s; server api-2:8000 max_fails=3 fail_timeout=30s; keepalive 32;}limit_req_zone $binary_remote_addr zone=perip:10m rate=20r/s;server { listen 443 ssl; http2 on; # nginx 1.25.1+; older configs put http2 on the listen line server_name api.acme.com; ssl_certificate /etc/ssl/fullchain.pem; ssl_certificate_key /etc/ssl/privkey.pem; ssl_protocols TLSv1.2 TLSv1.3; add_header Strict-Transport-Security "max-age=63072000; includeSubDomains" always; add_header X-Content-Type-Options nosniff always; add_header X-Frame-Options DENY always; server_tokens off; client_max_body_size 10m; client_body_timeout 10s; location /predict { limit_req zone=perip burst=40 nodelay; proxy_pass http://inference; proxy_read_timeout 30s; proxy_set_header X-Request-ID $request_id; proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for; }}least_conn rather than round-robin matters specifically for inference: request durations vary a lot (a small image versus a large one), and round-robin will hand a new request to a replica already processing three slow ones. Least-connections routes to whichever replica is actually free.
client_body_timeout 10s defends against slow-body attacks, where an attacker opens hundreds of connections and sends one byte per second, holding your workers indefinitely. The proxy absorbs that; your Python process never sees it.
Scaling horizontally
1apiVersion: apps/v12kind: Deployment3metadata: { name: inference-api }4spec:5 replicas: 36 template:7 spec:8 containers:9 - name: api10 image: ghcr.io/acme/inference:2.6.011 resources:12 requests: { cpu: "1000m", memory: "2Gi" }13 limits: { cpu: "2000m", memory: "4Gi" }14 env:15 - { name: OMP_NUM_THREADS, value: "1" }16 - name: JWT_SECRET17 valueFrom: { secretKeyRef: { name: api-secrets, key: jwt-secret } }18 livenessProbe:19 httpGet: { path: /healthz, port: 8000 }20 initialDelaySeconds: 4521 periodSeconds: 2022 readinessProbe:23 httpGet: { path: /readyz, port: 8000 }24 initialDelaySeconds: 2025 periodSeconds: 526 securityContext:27 runAsNonRoot: true28 runAsUser: 1000129 readOnlyRootFilesystem: true30 allowPrivilegeEscalation: false31 capabilities: { drop: ["ALL"] }32 terminationGracePeriodSeconds: 4533---34apiVersion: autoscaling/v235kind: HorizontalPodAutoscaler36metadata: { name: inference-api }37spec:38 scaleTargetRef: { apiVersion: apps/v1, kind: Deployment, name: inference-api }39 minReplicas: 340 maxReplicas: 2041 metrics:42 - type: Resource43 resource: { name: cpu, target: { type: Utilization, averageUtilization: 60 } }44 behavior:45 scaleUp: { stabilizationWindowSeconds: 30 }46 scaleDown: { stabilizationWindowSeconds: 300 }The autoscaler's arithmetic is worth knowing because it explains the behaviour you observe:
With 4 replicas averaging 85% CPU against a 60% target: ceil(4 × 85/60) = ceil(5.67) = 6 replicas. Note that the target of 60%, not 90%, is deliberate — queueing delay rises as roughly 1/(1 − utilisation), so running at 90% means latency five times higher than at 50%, and leaves no headroom for the two minutes a new pod takes to become ready.
terminationGracePeriodSeconds: 45 and a matching graceful shutdown in the app are what stop deployments from dropping requests. On SIGTERM, Kubernetes removes the pod from the service endpoints and signals the process; if your server exits immediately, every in-flight inference is lost. Uvicorn handles this correctly only if it receives the signal — which requires the exec form of CMD, since /bin/sh -c does not forward signals to its child.
Caching, circuit breakers, and retries
Caching
Cache only where inputs genuinely repeat. Hash the canonicalised input as the key, and always include the model version — otherwise a deployment serves stale predictions from the previous model.
1def cache_key(features: dict, model_version: str) -> str:2 canonical = json.dumps(features, sort_keys=True, separators=(",", ":"))3 return "pred:" + hashlib.sha256(f"{model_version}|{canonical}".encode()).hexdigest()45async def predict_cached(features: dict, version: str):6 key = cache_key(features, version)7 if (hit := await r.get(key)) is not None:8 CACHE_HITS.inc()9 return json.loads(hit)10 result = await run_model(features)11 await r.setex(key, 3600, json.dumps(result))12 return resultWork out whether it is worth it. With a 20% hit rate, a 2 ms cache read and a 90 ms inference:
A 19.6% latency reduction and, more usefully, 20% more capacity from the same instances. At a 3% hit rate the same arithmetic gives 87.4 ms — under 3% better, and not worth the operational cost of running and invalidating a cache. Measure your repeat rate before building this.
Circuit breakers
When a dependency — a feature store, a database — starts failing, retrying against it makes things worse and ties up your workers waiting on timeouts. A circuit breaker stops calling a broken dependency and fails fast instead.
1class CircuitBreaker:2 def __init__(self, threshold=5, recovery=30.0):3 self.threshold, self.recovery = threshold, recovery4 self.failures, self.opened_at, self.state = 0, None, "closed"56 async def call(self, fn, *args):7 if self.state == "open":8 if time.monotonic() - self.opened_at < self.recovery:9 raise HTTPException(503, "dependency unavailable (circuit open)")10 self.state = "half_open" # let ONE request through to test11 try:12 result = await fn(*args)13 except Exception:14 self.failures += 115 if self.failures >= self.threshold or self.state == "half_open":16 self.state, self.opened_at = "open", time.monotonic()17 raise18 self.failures, self.state = 0, "closed"19 return resultThe three states are the whole design. Closed: calls pass through, failures counted. Open: calls fail immediately without touching the dependency, giving it room to recover. Half-open: after the recovery window, exactly one request is allowed through — success closes the circuit, failure reopens it. Without the half-open state you either stay broken forever or slam the recovering dependency with full traffic the instant the timer expires.
Retries with backoff and jitter
1import random, asyncio23async def with_retry(fn, attempts=3, base=0.1, cap=2.0):4 for i in range(attempts):5 try:6 return await fn()7 except TransientError:8 if i == attempts - 1:9 raise10 delay = min(cap, base * (2 ** i))11 await asyncio.sleep(delay * (0.5 + random.random())) # jitter: 50-150% of delayThe jitter is the part people omit, and its absence causes a specific outage. Without it, a thousand clients that all failed at the same instant retry at exactly 0.1 s, then exactly 0.2 s, then exactly 0.4 s — synchronised thundering herds that keep the recovering service down. Randomising each delay spreads them out.
Retry only idempotent operations, and only on transient failures. Retrying a 400 is pointless; retrying a request that charges a card is dangerous.
Tracing, privacy, and being able to explain a decision
Distributed tracing
When one request touches a gateway, a feature store, a model, and a database, per-service latency numbers cannot tell you where the time went. Tracing joins the spans.
1from opentelemetry import trace2from opentelemetry.instrumentation.fastapi import FastAPIInstrumentor34FastAPIInstrumentor.instrument_app(app)5tracer = trace.get_tracer(__name__)67@app.post("/predict")8async def predict(req: PredictRequest):9 with tracer.start_as_current_span("fetch_features") as span:10 features = await feature_store.get(req.entity_id)11 span.set_attribute("feature.count", len(features))12 with tracer.start_as_current_span("inference") as span:13 span.set_attribute("model.version", MODEL_VERSION)14 result = await run_model(features)15 span.set_attribute("prediction.confidence", result["confidence"])16 return resultSet attributes that are bounded and useful — model version, feature count, confidence. Never set an attribute containing raw input data; traces are stored and searchable by anyone with access to the tracing system, which is usually a wider group than has access to your database.
Do not log what you do not need
1SENSITIVE = {"ssn", "email", "phone", "address", "account_number", "dob"}23def safe_for_logs(features: dict) -> dict:4 out = {}5 for k, v in features.items():6 if k in SENSITIVE:7 out[k] = "sha256:" + hashlib.sha256(str(v).encode()).hexdigest()[:12]8 elif isinstance(v, (int, float, bool)):9 out[k] = v10 else:11 out[k] = f"<{type(v).__name__} len={len(str(v))}>"12 return outHashing rather than dropping preserves what you actually need from a log: the ability to tell whether two requests concerned the same person, and to correlate a complaint with a specific prediction. Retention matters as much as content — logs kept for two years are a two-year liability, and under most data-protection regimes you must be able to delete a person's records on request, which is impossible if their data is smeared across raw log lines in cold storage.
Be able to explain a decision
For credit, hiring, insurance, and healthcare, a regulator or a customer can require the reason for an individual decision. Reconstructing it later is impossible unless you stored enough at the time.
1@app.post("/predict/explain")2async def predict_with_reasons(req: Applicant, claims = Depends(require_scope("explain"))):3 x = to_frame(req)4 score = float(model.predict_proba(x)[0, 1])5 contributions = explainer.shap_values(x)[0] # per-feature contribution6 top = sorted(zip(FEATURE_NAMES, contributions),7 key=lambda kv: abs(kv[1]), reverse=True)[:5]89 record = {10 "decision_id": uuid.uuid4().hex,11 "timestamp": utcnow().isoformat(),12 "model_version": MODEL_VERSION,13 "model_sha256": MODEL_DIGEST,14 "inputs": safe_for_logs(req.model_dump()),15 "score": round(score, 4),16 "threshold": DECISION_THRESHOLD,17 "outcome": "approve" if score >= DECISION_THRESHOLD else "decline",18 "top_factors": [{"feature": f, "contribution": round(float(c), 4)} for f, c in top],19 }20 await audit_store.write(record) # immutable, retained21 return {"decision_id": record["decision_id"], "outcome": record["outcome"],22 "reasons": [f for f, _ in top]}Storing model_sha256 alongside the version is what makes the record defensible. "Version 2.6.0" is a label someone could have reused; a digest identifies the exact bytes that produced the decision, so you can reload that artefact years later and reproduce the result.
An audit record written after the fact is a reconstruction; only one written at decision time is evidence.
What to do before your endpoint is reachable from outside
Most of the controls above are individually small. The reason services ship without them is that no single one feels urgent on the day. A short, ordered checklist fixes that — work down it and stop when the remaining items genuinely do not apply.
Authentication on every route, including the ones you added for debugging. Enumerate your routes programmatically and assert that each has an auth dependency; a forgotten /debug/predict is how "internal only" services end up public.
A hard cap on every list and every payload. Batch size, file size, string length. One unbounded field is one denial-of-service.
Rate limits in shared state, with a daily cap as well as a per-second one. The per-second limit protects capacity; the daily cap protects the model itself.
TLS terminated at a proxy, with body and header timeouts set there. Not in your Python process.
Containers running as a non-root user with a read-only root filesystem and all capabilities dropped. Four lines of YAML that remove an entire class of escalation.
A grep of your logs for personal data. Actually run it against a sample of production log lines rather than reasoning about what the code should be doing. Teams are routinely surprised.
Then run the exercise that surfaces what checklists miss: take one real request and follow it end to end, asking at each hop what an attacker who controlled the previous hop could do. That is where you find the debug route, the unbounded field, and the log line nobody meant to write.