package auth import ( "context" "crypto" "crypto/ecdsa" "crypto/ed25519" "crypto/elliptic" "crypto/rsa" "crypto/sha256" "crypto/sha512" "encoding/base64" "encoding/json" "errors" "fmt" "io" "math/big" "net/http" "net/url" "strings" "sync" "time" "github.com/example/ollama-fair-gateway/internal/config" ) type discoveryDoc struct { Issuer string `json:"issuer"` JWKSURI string `json:"jwks_uri"` AuthorizationEndpoint string `json:"authorization_endpoint"` TokenEndpoint string `json:"token_endpoint"` } type jwksDoc struct { Keys []json.RawMessage `json:"keys"` } type jwkHeader struct { Kty string `json:"kty"` Kid string `json:"kid"` Alg string `json:"alg"` Use string `json:"use"` N string `json:"n"` E string `json:"e"` Crv string `json:"crv"` X string `json:"x"` Y string `json:"y"` } type jwtHeader struct { Alg string `json:"alg"` Kid string `json:"kid"` Typ string `json:"typ"` } type OIDCVerifier struct { cfg config.OIDCConfig issuer string jwksURI string authorizationEndpoint string tokenEndpoint string client *http.Client mu sync.RWMutex keys map[string]crypto.PublicKey lastRefresh time.Time allowed map[string]bool adminGroups map[string]bool } func NewOIDCVerifier(ctx context.Context, cfg config.OIDCConfig) (*OIDCVerifier, error) { issuer := strings.TrimSuffix(cfg.Issuer, "/") v := &OIDCVerifier{cfg: cfg, issuer: issuer, client: &http.Client{Timeout: 10 * time.Second}, keys: map[string]crypto.PublicKey{}, allowed: map[string]bool{}, adminGroups: map[string]bool{}} for _, a := range cfg.AllowedAlgorithms { v.allowed[a] = true } for _, g := range cfg.AdminGroups { v.adminGroups[g] = true } req, err := http.NewRequestWithContext(ctx, http.MethodGet, issuer+"/.well-known/openid-configuration", nil) if err != nil { return nil, err } resp, err := v.client.Do(req) if err != nil { return nil, fmt.Errorf("oidc discovery: %w", err) } defer resp.Body.Close() if resp.StatusCode/100 != 2 { return nil, fmt.Errorf("oidc discovery: HTTP %d", resp.StatusCode) } var d discoveryDoc if err := json.NewDecoder(io.LimitReader(resp.Body, 1<<20)).Decode(&d); err != nil { return nil, err } if strings.TrimSuffix(d.Issuer, "/") != issuer { return nil, fmt.Errorf("oidc discovery issuer mismatch: got %q", d.Issuer) } if d.JWKSURI == "" { return nil, errors.New("oidc discovery returned no jwks_uri") } v.jwksURI = d.JWKSURI v.authorizationEndpoint = d.AuthorizationEndpoint v.tokenEndpoint = d.TokenEndpoint if err := v.refresh(ctx, true); err != nil { return nil, err } return v, nil } func (v *OIDCVerifier) Verify(ctx context.Context, token string) (Identity, error) { parts := strings.Split(token, ".") if len(parts) != 3 { return Identity{}, errors.New("token is not a JWT") } hb, err := base64.RawURLEncoding.DecodeString(parts[0]) if err != nil { return Identity{}, err } var h jwtHeader if err := json.Unmarshal(hb, &h); err != nil { return Identity{}, err } if h.Alg == "none" || !v.allowed[h.Alg] { return Identity{}, fmt.Errorf("JWT algorithm %q is not allowed", h.Alg) } key := v.key(h.Kid) if key == nil { if err := v.refresh(ctx, false); err != nil { return Identity{}, err } key = v.key(h.Kid) } if key == nil { return Identity{}, fmt.Errorf("no signing key for kid %q", h.Kid) } sig, err := base64.RawURLEncoding.DecodeString(parts[2]) if err != nil { return Identity{}, err } signingInput := []byte(parts[0] + "." + parts[1]) if err := verifySignature(h.Alg, key, signingInput, sig); err != nil { // Some providers rotate key material while reusing a kid. Refresh once // (subject to the anti-thundering-herd interval) before rejecting. _ = v.refresh(ctx, false) key = v.key(h.Kid) if key == nil { return Identity{}, err } if err2 := verifySignature(h.Alg, key, signingInput, sig); err2 != nil { return Identity{}, err2 } } pb, err := base64.RawURLEncoding.DecodeString(parts[1]) if err != nil { return Identity{}, err } dec := json.NewDecoder(strings.NewReader(string(pb))) dec.UseNumber() claims := map[string]any{} if err := dec.Decode(&claims); err != nil { return Identity{}, err } now := time.Now() skew := v.cfg.ClockSkew.Value() if iss, _ := claims["iss"].(string); strings.TrimSuffix(iss, "/") != v.issuer { return Identity{}, errors.New("issuer mismatch") } if !audContains(claims["aud"], v.cfg.Audience) { return Identity{}, errors.New("audience mismatch") } if exp, ok := numericTime(claims["exp"]); !ok || now.After(exp.Add(skew)) { return Identity{}, errors.New("token expired or missing exp") } if nbf, ok := numericTime(claims["nbf"]); ok && now.Add(skew).Before(nbf) { return Identity{}, errors.New("token not valid yet") } sub, _ := claims["sub"].(string) if sub == "" { return Identity{}, errors.New("token has no sub") } tenant := claimString(claims, v.cfg.TenantClaim) if tenant == "" { tenant = sub } app := claimString(claims, v.cfg.ApplicationClaim) scopes := map[string]bool{} for _, s := range claimStrings(claims["scope"]) { for _, p := range strings.Fields(s) { scopes[p] = true } } for _, s := range claimStrings(claims["scp"]) { for _, p := range strings.Fields(s) { scopes[p] = true } } groups := claimStrings(claimValue(claims, v.cfg.GroupsClaim)) for _, g := range groups { if v.adminGroups[g] { scopes["gateway:admin"] = true } } return Identity{Tenant: tenant, Subject: sub, Application: app, AuthType: "oidc", Scopes: scopes}, nil } func (v *OIDCVerifier) key(kid string) crypto.PublicKey { v.mu.RLock() defer v.mu.RUnlock() if kid != "" { return v.keys[kid] } if len(v.keys) == 1 { for _, k := range v.keys { return k } } return nil } func (v *OIDCVerifier) refresh(ctx context.Context, force bool) error { v.mu.Lock() defer v.mu.Unlock() if !force && time.Since(v.lastRefresh) < v.cfg.JWKSRefreshMinInterval.Value() { return nil } req, err := http.NewRequestWithContext(ctx, http.MethodGet, v.jwksURI, nil) if err != nil { return err } resp, err := v.client.Do(req) if err != nil { return fmt.Errorf("fetch jwks: %w", err) } defer resp.Body.Close() if resp.StatusCode/100 != 2 { return fmt.Errorf("fetch jwks: HTTP %d", resp.StatusCode) } var d jwksDoc if err := json.NewDecoder(io.LimitReader(resp.Body, 2<<20)).Decode(&d); err != nil { return err } keys := map[string]crypto.PublicKey{} for _, raw := range d.Keys { var j jwkHeader if err := json.Unmarshal(raw, &j); err != nil { continue } k, err := parseJWK(j) if err != nil { continue } keys[j.Kid] = k } if len(keys) == 0 { return errors.New("jwks contained no supported signing keys") } v.keys = keys v.lastRefresh = time.Now() return nil } func parseJWK(j jwkHeader) (crypto.PublicKey, error) { dec := base64.RawURLEncoding.DecodeString switch j.Kty { case "RSA": nB, err := dec(j.N) if err != nil { return nil, err } eB, err := dec(j.E) if err != nil { return nil, err } e := 0 for _, b := range eB { e = e<<8 + int(b) } if e == 0 { return nil, errors.New("bad RSA exponent") } return &rsa.PublicKey{N: new(big.Int).SetBytes(nB), E: e}, nil case "EC": xb, err := dec(j.X) if err != nil { return nil, err } yb, err := dec(j.Y) if err != nil { return nil, err } var c elliptic.Curve switch j.Crv { case "P-256": c = elliptic.P256() case "P-384": c = elliptic.P384() case "P-521": c = elliptic.P521() default: return nil, errors.New("unsupported EC curve") } x, y := new(big.Int).SetBytes(xb), new(big.Int).SetBytes(yb) if !c.IsOnCurve(x, y) { return nil, errors.New("EC key is not on curve") } return &ecdsa.PublicKey{Curve: c, X: x, Y: y}, nil case "OKP": if j.Crv != "Ed25519" { return nil, errors.New("unsupported OKP curve") } b, err := dec(j.X) if err != nil { return nil, err } if len(b) != ed25519.PublicKeySize { return nil, errors.New("bad Ed25519 key") } return ed25519.PublicKey(b), nil default: return nil, errors.New("unsupported key type") } } func verifySignature(alg string, key crypto.PublicKey, msg, sig []byte) error { var hash crypto.Hash var digest []byte switch alg { case "RS256", "PS256", "ES256": h := sha256.Sum256(msg) hash = crypto.SHA256 digest = h[:] case "RS384", "PS384", "ES384": h := sha512.Sum384(msg) hash = crypto.SHA384 digest = h[:] case "RS512", "PS512", "ES512": h := sha512.Sum512(msg) hash = crypto.SHA512 digest = h[:] case "EdDSA": pk, ok := key.(ed25519.PublicKey) if !ok || !ed25519.Verify(pk, msg, sig) { return errors.New("invalid JWT signature") } return nil default: return errors.New("unsupported JWT algorithm") } if strings.HasPrefix(alg, "RS") { pk, ok := key.(*rsa.PublicKey) if !ok { return errors.New("JWT key type mismatch") } if err := rsa.VerifyPKCS1v15(pk, hash, digest, sig); err != nil { return errors.New("invalid JWT signature") } return nil } if strings.HasPrefix(alg, "PS") { pk, ok := key.(*rsa.PublicKey) if !ok { return errors.New("JWT key type mismatch") } if err := rsa.VerifyPSS(pk, hash, digest, sig, &rsa.PSSOptions{SaltLength: rsa.PSSSaltLengthEqualsHash, Hash: hash}); err != nil { return errors.New("invalid JWT signature") } return nil } pk, ok := key.(*ecdsa.PublicKey) if !ok { return errors.New("JWT key type mismatch") } n := (pk.Curve.Params().BitSize + 7) / 8 if len(sig) != 2*n { return errors.New("bad ECDSA JWT signature length") } r, s := new(big.Int).SetBytes(sig[:n]), new(big.Int).SetBytes(sig[n:]) if !ecdsa.Verify(pk, digest, r, s) { return errors.New("invalid JWT signature") } return nil } func audContains(v any, want string) bool { for _, s := range claimStrings(v) { if s == want { return true } } return false } func claimStrings(v any) []string { switch x := v.(type) { case string: return []string{x} case []any: out := []string{} for _, z := range x { if s, ok := z.(string); ok { out = append(out, s) } } return out case []string: return x } return nil } func numericTime(v any) (time.Time, bool) { var n int64 switch x := v.(type) { case json.Number: i, e := x.Int64() if e != nil { return time.Time{}, false } n = i case float64: n = int64(x) case int64: n = x default: return time.Time{}, false } return time.Unix(n, 0), true } func claimValue(m map[string]any, path string) any { if path == "" { return nil } var cur any = m for _, p := range strings.Split(path, ".") { mm, ok := cur.(map[string]any) if !ok { return nil } cur = mm[p] } return cur } func claimString(m map[string]any, path string) string { if s, ok := claimValue(m, path).(string); ok { return s } return "" } type BrowserEndpoints struct { Authorization string `json:"authorization_endpoint"` Token string `json:"token_endpoint"` } type TokenExchange struct { AccessToken string `json:"access_token"` TokenType string `json:"token_type"` ExpiresIn int64 `json:"expires_in"` IDToken string `json:"id_token,omitempty"` Scope string `json:"scope,omitempty"` } func (v *OIDCVerifier) BrowserEndpoints() BrowserEndpoints { return BrowserEndpoints{Authorization: v.authorizationEndpoint, Token: v.tokenEndpoint} } func (v *OIDCVerifier) ExchangeCode(ctx context.Context, code, redirectURI, clientID, clientSecret, verifier string) (TokenExchange, error) { if v.tokenEndpoint == "" { return TokenExchange{}, errors.New("OIDC discovery returned no token_endpoint") } form := url.Values{} form.Set("grant_type", "authorization_code") form.Set("code", code) form.Set("redirect_uri", redirectURI) form.Set("client_id", clientID) if verifier != "" { form.Set("code_verifier", verifier) } if clientSecret != "" { form.Set("client_secret", clientSecret) } req, err := http.NewRequestWithContext(ctx, http.MethodPost, v.tokenEndpoint, strings.NewReader(form.Encode())) if err != nil { return TokenExchange{}, err } req.Header.Set("Content-Type", "application/x-www-form-urlencoded") req.Header.Set("Accept", "application/json") resp, err := v.client.Do(req) if err != nil { return TokenExchange{}, fmt.Errorf("OIDC token exchange: %w", err) } defer resp.Body.Close() body, err := io.ReadAll(io.LimitReader(resp.Body, 2<<20)) if err != nil { return TokenExchange{}, err } if resp.StatusCode/100 != 2 { return TokenExchange{}, fmt.Errorf("OIDC token exchange: HTTP %d", resp.StatusCode) } var out TokenExchange if err := json.Unmarshal(body, &out); err != nil { return TokenExchange{}, err } if out.AccessToken == "" { return TokenExchange{}, errors.New("OIDC token response contains no access_token") } return out, nil }