From b0ab0ce92f5d92bad1de45440d02e6b1ddbe0788 Mon Sep 17 00:00:00 2001 From: hanzo-dev Date: Sat, 4 Jul 2026 11:28:21 -0700 Subject: [PATCH] fix(auth): kid-first JWT verification, drop per-attempt debug logging MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit validateToken logged one line per JWKS key per request (a leftover DEBUG line plus a per-failed-key line) — with multi-brand certs that produced a log line per brand cert per request, ~4.4M lines/12h in production, and brute-forced every key on every call. Now: verify against the key the token's kid header names (single signature check in the common case), fall back across the remaining keys only on rotation skew, and emit no per-attempt logs — the returned error carries the diagnosis. --- cmd/kms/auth.go | 30 +++++++++++++++++++++++------- 1 file changed, 23 insertions(+), 7 deletions(-) diff --git a/cmd/kms/auth.go b/cmd/kms/auth.go index f6b24e4f1..0e6f49114 100644 --- a/cmd/kms/auth.go +++ b/cmd/kms/auth.go @@ -149,28 +149,44 @@ func (a *orgJWTAuth) validate(ctx context.Context, raw string) (*orgClaims, erro if err != nil { return nil, fmt.Errorf("jwks: %w", err) } - // DEBUG: log JWKS state + token kid for diagnosis tokKid := "" if len(tok.Headers) > 0 { tokKid = tok.Headers[0].KeyID } - kidList := make([]string, 0, len(keys.Keys)) - for _, k := range keys.Keys { - kidList = append(kidList, k.KeyID) - } - log.Printf("kms: validate: tok kid=%s jwks kids=%v", tokKid, kidList) var claims orgClaims verified := false var lastErr error + // Try the key the token names first (kid header) — the common case is one + // exact match, so verification is a single signature check instead of + // brute-forcing every brand cert. A miss on a non-matching key is expected, + // not an event: no per-attempt logging (it produced one line per brand cert + // per request — millions/day); the returned error carries the diagnosis. for _, k := range keys.Keys { + if tokKid != "" && k.KeyID != tokKid { + continue + } if err := tok.Claims(k.Key, &claims); err == nil { verified = true break } else { - log.Printf("kms: validate: kid=%s err=%v", k.KeyID, err) lastErr = err } } + if !verified && tokKid != "" { + // kid named but didn't verify (rotation skew) — legacy fallback across + // the remaining keys before failing. + for _, k := range keys.Keys { + if k.KeyID == tokKid { + continue + } + if err := tok.Claims(k.Key, &claims); err == nil { + verified = true + break + } else { + lastErr = err + } + } + } if !verified { if lastErr == nil { lastErr = errors.New("no key matched")