Model Deployment for AI Engineers

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.

Each layer rejects what the next would have paid forEdge — TLS termination and a WAFAuth — API key or a verified JWT claimRate limit — shared counterin Redis, not per replicaValidation — schema, size and type before decodingModel replicas — the only expensive layer
A private subnet is not authentication; without the auth layer, anyone inside can extract the model by querying it.

The threat model

A model-serving API inherits every web-API threat and adds several of its own.

ThreatWhat the attacker doesPrimary control
Unauthorised accessCalls the endpoint directlyAuthentication on every route
Model extractionSystematic queries to clone the modelRate limits per identity, query-pattern detection, coarse outputs
Membership inferenceDetermines whether a specific record was in your training dataAvoid returning raw confidences; differential privacy in training
Adversarial inputCrafts an input that forces a chosen wrong outputInput validation, ensembling, adversarial training
Resource exhaustionHuge payloads or unbounded batchesSize caps, batch caps, timeouts, quotas
Data exfiltration via logsReads personal data your service loggedLog hashes and shapes, never raw inputs
Prompt injection (generative models)Embeds instructions in user text that the model obeysSeparate 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 10510^{5} 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.

Python
import hashlib, hmac, secretsfrom fastapi import Security, HTTPException, statusfrom fastapi.security import APIKeyHeaderapi_key_header = APIKeyHeader(name="X-API-Key", auto_error=False)def issue_key() -> tuple[str, str]:    raw = "sk_" + secrets.token_urlsafe(32)          # shown to the user exactly once    return raw, hashlib.sha256(raw.encode()).hexdigest()async def require_key(key: str = Security(api_key_header)):    if not key:        raise HTTPException(status.HTTP_401_UNAUTHORIZED, "missing X-API-Key")    digest = hashlib.sha256(key.encode()).hexdigest()    record = await db.fetch_one(        "SELECT * FROM api_keys WHERE key_hash = :h AND revoked_at IS NULL",        {"h": digest})    if record is None or not hmac.compare_digest(record["key_hash"], digest):        raise HTTPException(status.HTTP_401_UNAUTHORIZED, "invalid key")    if record["expires_at"] < utcnow():        raise HTTPException(status.HTTP_401_UNAUTHORIZED, "key expired")    return record@app.post("/predict")async def predict(req: PredictRequest, client = Depends(require_key)):    ...

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.

Python
import jwtfrom datetime import datetime, timedelta, timezonefrom fastapi.security import HTTPBearer, HTTPAuthorizationCredentialsSECRET, ALGO = os.environ["JWT_SECRET"], "HS256"bearer = HTTPBearer()def create_token(user_id: str, scopes: list[str], minutes: int = 30) -> str:    now = datetime.now(timezone.utc)    return jwt.encode({        "sub": user_id, "scopes": scopes,        "iat": now, "exp": now + timedelta(minutes=minutes),        "iss": "auth.acme.com", "aud": "inference-api",    }, SECRET, algorithm=ALGO)def require_scope(scope: str):                # plain def: called at import time    async def dependency(cred: HTTPAuthorizationCredentials = Security(bearer)):        try:            claims = jwt.decode(cred.credentials, SECRET,                                algorithms=[ALGO],          # never accept "none"                                audience="inference-api",                                issuer="auth.acme.com")        except jwt.ExpiredSignatureError:            raise HTTPException(401, "token expired")        except jwt.InvalidTokenError as e:            raise HTTPException(401, f"invalid token: {e}")        if scope not in claims.get("scopes", []):            raise HTTPException(403, f"missing scope: {scope}")        return claims    return dependency@app.post("/predict")async def predict(req: PredictRequest, claims = Depends(require_scope("predict:write"))):    ...

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 keyJWTmTLS
CallerAnother serviceAn end user via an appAnother service, high trust
Carries permissionsVia a database lookupIn the token itselfIn the certificate
RevocationImmediateOnly at expiry, or via a denylistVia CRL / short-lived certs
Per-request costOne DB or cache lookupSignature verification onlyTLS handshake, then free
Main riskKey committed to a repositoryLong expiry, unpinned algorithmCertificate 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.

Python
from pydantic import BaseModel, Field, field_validatorfrom typing import Literal, Listimport mathclass Applicant(BaseModel):    annual_income: float = Field(..., ge=0, le=10_000_000)    age: int = Field(..., ge=18, le=120)    debt_ratio: float = Field(..., ge=0, le=5.0)    employment: Literal["full_time", "part_time", "self_employed", "retired"]    @field_validator("annual_income", "debt_ratio")    @classmethod    def finite(cls, v: float) -> float:        if not math.isfinite(v):                # NaN and inf pass float() happily            raise ValueError("value must be finite")        return vclass BatchRequest(BaseModel):    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.

Python
import redis.asyncio as redisr = redis.from_url(os.environ["REDIS_URL"])# Atomic sliding window. Executed server-side so concurrent calls cannot interleave.SLIDING_WINDOW = """local key, now, window, limit = KEYS[1], tonumber(ARGV[1]), tonumber(ARGV[2]), tonumber(ARGV[3])redis.call('ZREMRANGEBYSCORE', key, 0, now - window)local used = redis.call('ZCARD', key)if used < limit then  redis.call('ZADD', key, now, ARGV[4])  redis.call('EXPIRE', key, window)  return {1, limit - used - 1}endreturn {0, 0}"""_script = r.register_script(SLIDING_WINDOW)TIERS = {"free": (100, 3600), "pro": (10_000, 3600), "enterprise": (200_000, 3600)}async def enforce_quota(client = Depends(require_key)):    limit, window = TIERS[client["tier"]]    allowed, remaining = await _script(        keys=[f"rl:{client['id']}"],        args=[time.time(), window, limit, uuid.uuid4().hex])    if not allowed:        raise HTTPException(429, "rate limit exceeded",                            headers={"Retry-After": str(window)})    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.

Text
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

YAML
apiVersion: apps/v1kind: Deploymentmetadata: { name: inference-api }spec:  replicas: 3  template:    spec:      containers:      - name: api        image: ghcr.io/acme/inference:2.6.0        resources:          requests: { cpu: "1000m", memory: "2Gi" }          limits:   { cpu: "2000m", memory: "4Gi" }        env:        - { name: OMP_NUM_THREADS, value: "1" }        - name: JWT_SECRET          valueFrom: { secretKeyRef: { name: api-secrets, key: jwt-secret } }        livenessProbe:          httpGet: { path: /healthz, port: 8000 }          initialDelaySeconds: 45          periodSeconds: 20        readinessProbe:          httpGet: { path: /readyz, port: 8000 }          initialDelaySeconds: 20          periodSeconds: 5        securityContext:          runAsNonRoot: true          runAsUser: 10001          readOnlyRootFilesystem: true          allowPrivilegeEscalation: false          capabilities: { drop: ["ALL"] }      terminationGracePeriodSeconds: 45---apiVersion: autoscaling/v2kind: HorizontalPodAutoscalermetadata: { name: inference-api }spec:  scaleTargetRef: { apiVersion: apps/v1, kind: Deployment, name: inference-api }  minReplicas: 3  maxReplicas: 20  metrics:  - type: Resource    resource: { name: cpu, target: { type: Utilization, averageUtilization: 60 } }  behavior:    scaleUp:   { stabilizationWindowSeconds: 30 }    scaleDown: { stabilizationWindowSeconds: 300 }

The autoscaler's arithmetic is worth knowing because it explains the behaviour you observe:

desired=⌈current×current metrictarget metric⌉\text{desired} = \left\lceil \text{current} \times \frac{\text{current metric}}{\text{target metric}} \right\rceil

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.

Python
def cache_key(features: dict, model_version: str) -> str:    canonical = json.dumps(features, sort_keys=True, separators=(",", ":"))    return "pred:" + hashlib.sha256(f"{model_version}|{canonical}".encode()).hexdigest()async def predict_cached(features: dict, version: str):    key = cache_key(features, version)    if (hit := await r.get(key)) is not None:        CACHE_HITS.inc()        return json.loads(hit)    result = await run_model(features)    await r.setex(key, 3600, json.dumps(result))    return result

Work out whether it is worth it. With a 20% hit rate, a 2 ms cache read and a 90 ms inference:

0.20×2+0.80×90=0.4+72=72.4 ms mean, versus 90 ms0.20 \times 2 + 0.80 \times 90 = 0.4 + 72 = 72.4\ \text{ms mean, versus } 90\ \text{ms}

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.

Python
class CircuitBreaker:    def __init__(self, threshold=5, recovery=30.0):        self.threshold, self.recovery = threshold, recovery        self.failures, self.opened_at, self.state = 0, None, "closed"    async def call(self, fn, *args):        if self.state == "open":            if time.monotonic() - self.opened_at < self.recovery:                raise HTTPException(503, "dependency unavailable (circuit open)")            self.state = "half_open"          # let ONE request through to test        try:            result = await fn(*args)        except Exception:            self.failures += 1            if self.failures >= self.threshold or self.state == "half_open":                self.state, self.opened_at = "open", time.monotonic()            raise        self.failures, self.state = 0, "closed"        return result

The 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

Python
import random, asyncioasync def with_retry(fn, attempts=3, base=0.1, cap=2.0):    for i in range(attempts):        try:            return await fn()        except TransientError:            if i == attempts - 1:                raise            delay = min(cap, base * (2 ** i))            await asyncio.sleep(delay * (0.5 + random.random()))   # jitter: 50-150% of delay

The 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.

Python
from opentelemetry import tracefrom opentelemetry.instrumentation.fastapi import FastAPIInstrumentorFastAPIInstrumentor.instrument_app(app)tracer = trace.get_tracer(__name__)@app.post("/predict")async def predict(req: PredictRequest):    with tracer.start_as_current_span("fetch_features") as span:        features = await feature_store.get(req.entity_id)        span.set_attribute("feature.count", len(features))    with tracer.start_as_current_span("inference") as span:        span.set_attribute("model.version", MODEL_VERSION)        result = await run_model(features)        span.set_attribute("prediction.confidence", result["confidence"])    return result

Set 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

Python
SENSITIVE = {"ssn", "email", "phone", "address", "account_number", "dob"}def safe_for_logs(features: dict) -> dict:    out = {}    for k, v in features.items():        if k in SENSITIVE:            out[k] = "sha256:" + hashlib.sha256(str(v).encode()).hexdigest()[:12]        elif isinstance(v, (int, float, bool)):            out[k] = v        else:            out[k] = f"<{type(v).__name__} len={len(str(v))}>"    return out

Hashing 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.

Python
@app.post("/predict/explain")async def predict_with_reasons(req: Applicant, claims = Depends(require_scope("explain"))):    x = to_frame(req)    score = float(model.predict_proba(x)[0, 1])    contributions = explainer.shap_values(x)[0]              # per-feature contribution    top = sorted(zip(FEATURE_NAMES, contributions),                 key=lambda kv: abs(kv[1]), reverse=True)[:5]    record = {        "decision_id": uuid.uuid4().hex,        "timestamp": utcnow().isoformat(),        "model_version": MODEL_VERSION,        "model_sha256": MODEL_DIGEST,        "inputs": safe_for_logs(req.model_dump()),        "score": round(score, 4),        "threshold": DECISION_THRESHOLD,        "outcome": "approve" if score >= DECISION_THRESHOLD else "decline",        "top_factors": [{"feature": f, "contribution": round(float(c), 4)} for f, c in top],    }    await audit_store.write(record)                          # immutable, retained    return {"decision_id": record["decision_id"], "outcome": record["outcome"],            "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.