Files
ResolveSpec/pkg/security/oauth_jwk.go
T
Hein 640faeeeaf feat(security): full OAuth 2.1 / OpenID Connect server and OIDC relying-party client
Authorization server: consent and scopes, OIDC (nonce, auth_time, acr, sid,
at_hash, signed userinfo, RP-initiated and back-channel logout), managed
refresh tokens with rotation and reuse detection, RFC 9068 JWT access tokens,
DPoP, PAR, device grant, token exchange, private_key_jwt, RFC 7591/7592
registration, RFC 9207 iss, signing keyring with rotation.

State is DB-backed through a new lookup.OAuthGrantStore (procedure and direct
backends, four dialect DDLs, conformance cases).

Client side: WithOIDC discovery, PKCE, nonce, id_token validation, OAuth2LogoutURL.

PeekRefresh now returns already rotated tokens so RotateRefresh can detect reuse.

Docs: OAUTH2_SERVER.md, oauth2_full_example.go, breaking_changes.md step 8.
2026-10-01 14:42:12 +02:00

422 lines
11 KiB
Go

package security
import (
"context"
"crypto"
"crypto/ecdsa"
"crypto/elliptic"
"crypto/rand"
"crypto/rsa"
"crypto/sha256"
"encoding/base64"
"encoding/json"
"fmt"
"io"
"math/big"
"net"
"net/http"
"sort"
"sync"
"syscall"
"time"
"github.com/golang-jwt/jwt/v5"
)
// OAuthSigningKey is a key the authorization server signs tokens with. The first configured
// key is the default; the others are published in the JWKS so tokens signed before a rotation
// stay verifiable, and a client can ask for one by id_token_signed_response_alg.
type OAuthSigningKey struct {
// ID is the JWKS "kid". Derived from the public key (RFC 7638 thumbprint) when empty.
ID string
// Key is an *rsa.PrivateKey (RS256) or an *ecdsa.PrivateKey (ES256 for P-256, ES384 for P-384).
Key crypto.Signer
// Alg overrides the algorithm inferred from Key (RS256, PS256, ES256, ES384).
Alg string
}
type oauthKey struct {
id string
alg string
signer crypto.Signer
method jwt.SigningMethod
}
// oauthKeyring holds the server's signing keys.
type oauthKeyring struct {
keys []oauthKey
}
func newOAuthKeyring(cfg *OAuthServerConfig) (*oauthKeyring, error) {
in := cfg.SigningKeys
if len(in) == 0 && cfg.SigningKey != nil {
in = []OAuthSigningKey{{Key: cfg.SigningKey}}
}
if len(in) == 0 {
k, err := rsa.GenerateKey(rand.Reader, 2048)
if err != nil {
return nil, fmt.Errorf("generate signing key: %w", err)
}
in = []OAuthSigningKey{{Key: k}}
}
kr := &oauthKeyring{}
for _, k := range in {
if k.Key == nil {
return nil, fmt.Errorf("signing key without a private key")
}
alg := k.Alg
if alg == "" {
switch key := k.Key.(type) {
case *rsa.PrivateKey:
alg = "RS256"
case *ecdsa.PrivateKey:
switch key.Curve {
case elliptic.P256():
alg = "ES256"
case elliptic.P384():
alg = "ES384"
default:
return nil, fmt.Errorf("unsupported signing curve %s", key.Curve.Params().Name)
}
default:
return nil, fmt.Errorf("unsupported signing key type %T", k.Key)
}
}
method := jwt.GetSigningMethod(alg)
if method == nil {
return nil, fmt.Errorf("unsupported signing algorithm %q", alg)
}
id := k.ID
if id == "" {
tp, err := jwkThumbprint(k.Key.Public())
if err != nil {
return nil, err
}
id = tp[:16]
}
kr.keys = append(kr.keys, oauthKey{id: id, alg: alg, signer: k.Key, method: method})
}
return kr, nil
}
// forAlg returns the first key signing with alg, or the default key when alg is empty or unknown.
func (kr *oauthKeyring) forAlg(alg string) *oauthKey {
if alg != "" {
for i := range kr.keys {
if kr.keys[i].alg == alg {
return &kr.keys[i]
}
}
}
return &kr.keys[0]
}
func (kr *oauthKeyring) algs() []string {
var out []string
for _, k := range kr.keys {
if !oauthSliceContains(out, k.alg) {
out = append(out, k.alg)
}
}
return out
}
func (kr *oauthKeyring) jwks() []map[string]any {
out := make([]map[string]any, 0, len(kr.keys))
for _, k := range kr.keys {
if jwk, err := jwkFromPublic(k.signer.Public(), k.id, k.alg); err == nil {
out = append(out, jwk)
}
}
return out
}
// publicFor returns the public key with kid (or the default one when kid is empty).
func (kr *oauthKeyring) publicFor(kid string) (crypto.PublicKey, *oauthKey) {
for i := range kr.keys {
if kid == "" || kr.keys[i].id == kid {
return kr.keys[i].signer.Public(), &kr.keys[i]
}
}
return nil, nil
}
// sign signs claims with k, adding typ when non-empty.
func (k *oauthKey) sign(claims jwt.Claims, typ string) (string, error) {
t := jwt.NewWithClaims(k.method, claims)
t.Header["kid"] = k.id
if typ != "" {
t.Header["typ"] = typ
}
return t.SignedString(k.signer)
}
// --- JWK encoding ----------------------------------------------------------------------------
func b64u(b []byte) string { return base64.RawURLEncoding.EncodeToString(b) }
func jwkFromPublic(pub crypto.PublicKey, kid, alg string) (map[string]any, error) {
jwk := map[string]any{"use": "sig"}
if kid != "" {
jwk["kid"] = kid
}
if alg != "" {
jwk["alg"] = alg
}
switch k := pub.(type) {
case *rsa.PublicKey:
jwk["kty"] = "RSA"
jwk["n"] = b64u(k.N.Bytes())
jwk["e"] = b64u(big.NewInt(int64(k.E)).Bytes())
case *ecdsa.PublicKey:
size := (k.Curve.Params().BitSize + 7) / 8
jwk["kty"] = "EC"
jwk["crv"] = k.Curve.Params().Name
jwk["x"] = b64u(k.X.FillBytes(make([]byte, size)))
jwk["y"] = b64u(k.Y.FillBytes(make([]byte, size)))
default:
return nil, fmt.Errorf("unsupported key type %T", pub)
}
return jwk, nil
}
// jwkThumbprint is the RFC 7638 SHA-256 thumbprint (base64url).
func jwkThumbprint(pub crypto.PublicKey) (string, error) {
var canonical string
switch k := pub.(type) {
case *rsa.PublicKey:
canonical = fmt.Sprintf(`{"e":%q,"kty":"RSA","n":%q}`, b64u(big.NewInt(int64(k.E)).Bytes()), b64u(k.N.Bytes()))
case *ecdsa.PublicKey:
size := (k.Curve.Params().BitSize + 7) / 8
canonical = fmt.Sprintf(`{"crv":%q,"kty":"EC","x":%q,"y":%q}`, k.Curve.Params().Name,
b64u(k.X.FillBytes(make([]byte, size))), b64u(k.Y.FillBytes(make([]byte, size))))
default:
return "", fmt.Errorf("unsupported key type %T", pub)
}
sum := sha256.Sum256([]byte(canonical))
return b64u(sum[:]), nil
}
func publicFromJWK(m map[string]any) (crypto.PublicKey, error) {
str := func(k string) string { s, _ := m[k].(string); return s }
dec := func(k string) (*big.Int, error) {
raw, err := base64.RawURLEncoding.DecodeString(str(k))
if err != nil || len(raw) == 0 {
return nil, fmt.Errorf("invalid JWK member %q", k)
}
return new(big.Int).SetBytes(raw), nil
}
switch str("kty") {
case "RSA":
n, err := dec("n")
if err != nil {
return nil, err
}
e, err := dec("e")
if err != nil {
return nil, err
}
if n.BitLen() < 2048 || !e.IsInt64() || e.Int64() < 3 {
return nil, fmt.Errorf("RSA key too weak")
}
return &rsa.PublicKey{N: n, E: int(e.Int64())}, nil
case "EC":
var curve elliptic.Curve
switch str("crv") {
case "P-256":
curve = elliptic.P256()
case "P-384":
curve = elliptic.P384()
default:
return nil, fmt.Errorf("unsupported curve %q", str("crv"))
}
x, err := dec("x")
if err != nil {
return nil, err
}
y, err := dec("y")
if err != nil {
return nil, err
}
pub := &ecdsa.PublicKey{Curve: curve, X: x, Y: y}
if _, err := pub.ECDH(); err != nil { // rejects points that are not on the curve
return nil, fmt.Errorf("invalid EC point: %w", err)
}
return pub, nil
}
return nil, fmt.Errorf("unsupported key type %q", str("kty"))
}
// jwkEntry is one verification key of a JWK set.
type jwkEntry struct {
kid string
alg string
use string
pub crypto.PublicKey
}
// jwkSet is a parsed JWK set.
type jwkSet struct{ keys []jwkEntry }
func parseJWKS(data []byte) (*jwkSet, error) {
var doc struct {
Keys []map[string]any `json:"keys"`
}
if err := json.Unmarshal(data, &doc); err != nil {
return nil, fmt.Errorf("invalid JWKS: %w", err)
}
set := &jwkSet{}
for _, k := range doc.Keys {
pub, err := publicFromJWK(k)
if err != nil {
continue // skip keys of a type we cannot use
}
e := jwkEntry{pub: pub}
e.kid, _ = k["kid"].(string)
e.alg, _ = k["alg"].(string)
e.use, _ = k["use"].(string)
set.keys = append(set.keys, e)
}
if len(set.keys) == 0 {
return nil, fmt.Errorf("JWKS contains no usable keys")
}
return set, nil
}
// candidates returns the keys that may have produced a token with kid and alg.
func (s *jwkSet) candidates(kid, alg string) []crypto.PublicKey {
var out []crypto.PublicKey
for _, k := range s.keys {
if k.use != "" && k.use != "sig" {
continue
}
if kid != "" && k.kid != "" && k.kid != kid {
continue
}
if k.alg != "" && alg != "" && k.alg != alg {
continue
}
out = append(out, k.pub)
}
return out
}
// verifyJWTWithSet verifies token against the set. allowed lists the accepted algorithms.
func verifyJWTWithSet(token string, set *jwkSet, allowed []string, claims jwt.Claims, opts ...jwt.ParserOption) (*jwt.Token, error) {
opts = append([]jwt.ParserOption{jwt.WithValidMethods(allowed)}, opts...)
parser := jwt.NewParser(opts...)
lastErr := fmt.Errorf("no matching key")
unverified, _, err := parser.ParseUnverified(token, jwt.MapClaims{})
if err != nil {
return nil, err
}
kid, _ := unverified.Header["kid"].(string)
alg, _ := unverified.Header["alg"].(string)
for _, pub := range set.candidates(kid, alg) {
tok, err := parser.ParseWithClaims(token, claims, func(*jwt.Token) (any, error) { return pub, nil })
if err == nil {
return tok, nil
}
lastErr = err
}
return nil, lastErr
}
// --- remote JWKS fetching --------------------------------------------------------------------
// jwksCache fetches and caches remote JWK sets.
type jwksCache struct {
client *http.Client
ttl time.Duration
mu sync.Mutex
entries map[string]jwksCacheEntry
}
type jwksCacheEntry struct {
set *jwkSet
fetched time.Time
}
func newJWKSCache(client *http.Client) *jwksCache {
return &jwksCache{client: client, ttl: time.Hour, entries: map[string]jwksCacheEntry{}}
}
// get returns the set at uri. refresh forces a fetch (used when a kid is not in the cached set),
// rate limited to one fetch per 30 seconds per URI.
func (c *jwksCache) get(ctx context.Context, uri string, refresh bool) (*jwkSet, error) {
c.mu.Lock()
e, ok := c.entries[uri]
c.mu.Unlock()
age := time.Since(e.fetched)
if ok && age < c.ttl && (!refresh || age < 30*time.Second) {
return e.set, nil
}
req, err := http.NewRequestWithContext(ctx, http.MethodGet, uri, nil)
if err != nil {
return nil, err
}
resp, err := c.client.Do(req)
if err != nil {
if ok {
return e.set, nil // keep using the stale set while the endpoint is down
}
return nil, fmt.Errorf("fetch JWKS: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("fetch JWKS: status %d", resp.StatusCode)
}
body, err := io.ReadAll(io.LimitReader(resp.Body, 1<<20))
if err != nil {
return nil, err
}
set, err := parseJWKS(body)
if err != nil {
return nil, err
}
c.mu.Lock()
c.entries[uri] = jwksCacheEntry{set: set, fetched: time.Now()}
c.mu.Unlock()
return set, nil
}
// publicHTTPClient returns an HTTP client for fetching URLs supplied by clients. It refuses to
// connect to loopback, private and link-local addresses (SSRF) unless allowPrivate is set.
func publicHTTPClient(allowPrivate bool) *http.Client {
dialer := &net.Dialer{Timeout: 5 * time.Second}
if !allowPrivate {
dialer.Control = func(_, address string, _ syscall.RawConn) error {
host, _, err := net.SplitHostPort(address)
if err != nil {
return err
}
ip := net.ParseIP(host)
if ip == nil || ip.IsLoopback() || ip.IsPrivate() || ip.IsLinkLocalUnicast() || ip.IsLinkLocalMulticast() ||
ip.IsUnspecified() || ip.IsMulticast() {
return fmt.Errorf("address %s is not allowed", host)
}
return nil
}
}
return &http.Client{
Timeout: 10 * time.Second,
Transport: &http.Transport{DialContext: dialer.DialContext, Proxy: nil},
CheckRedirect: func(_ *http.Request, via []*http.Request) error {
if len(via) >= 3 {
return fmt.Errorf("too many redirects")
}
return nil
},
}
}
func sortedKeys(m map[string]struct{}) []string {
out := make([]string, 0, len(m))
for k := range m {
out = append(out, k)
}
sort.Strings(out)
return out
}