package oidc import ( "crypto/ecdsa" "crypto/elliptic" "crypto/rsa" "encoding/base64" "encoding/json" "fmt" "io" "math/big" "net/http" "strings" "sync" "time" "github.com/golang-jwt/jwt/v5" ) // jwk est une clé du document JWKS du fournisseur. type jwk struct { Kid string `json:"kid"` Kty string `json:"kty"` 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 jwksDocument struct { Keys []jwk `json:"keys"` } type jwksCacheEntry struct { keys map[string]interface{} expiresAt time.Time } // jwksStore maintient un cache des clés publiques de signature par fournisseur. type jwksStore struct { httpClient *http.Client mu sync.Mutex entries map[string]*jwksCacheEntry } func newJwksStore() *jwksStore { return &jwksStore{ httpClient: &http.Client{Timeout: 10 * time.Second}, entries: make(map[string]*jwksCacheEntry), } } // KeyFunc retourne une fonction de résolution de clé pour golang-jwt, // valable pour un fournisseur donné. func (s *jwksStore) KeyFunc(providerID string, jwksURI string) jwt.Keyfunc { return func(t *jwt.Token) (interface{}, error) { kid, ok := t.Header["kid"].(string) if !ok || kid == "" { return nil, fmt.Errorf("id token: kid manquant") } key, err := s.key(providerID, jwksURI, kid) if err != nil { return nil, err } return key, nil } } // key retourne la clé publique correspondant au kid, avec mise en cache. // Si le kid est inconnu, le cache est invalidé et le JWKS re-téléchargé // (rotation de clés côté fournisseur). func (s *jwksStore) key(providerID string, jwksURI string, kid string) (interface{}, error) { fetch := func() (map[string]interface{}, error) { keys, err := s.fetch(jwksURI) if err != nil { return nil, fmt.Errorf("error fetching provider JWKS: %s", err) } s.mu.Lock() s.entries[providerID] = &jwksCacheEntry{ keys: keys, expiresAt: time.Now().Add(1 * time.Hour), } s.mu.Unlock() return keys, nil } s.mu.Lock() entry, ok := s.entries[providerID] valid := ok && time.Now().Before(entry.expiresAt) s.mu.Unlock() if valid { if key, ok := entry.keys[kid]; ok { return key, nil } } keys, err := fetch() if err != nil { return nil, err } key, ok := keys[kid] if !ok { return nil, fmt.Errorf("id token: aucune clé JWKS ne correspond au kid %s", kid) } return key, nil } func (s *jwksStore) fetch(jwksURI string) (map[string]interface{}, error) { res, err := s.httpClient.Get(jwksURI) if err != nil { return nil, err } defer res.Body.Close() if res.StatusCode != http.StatusOK { return nil, fmt.Errorf("unexpected status %d", res.StatusCode) } body, err := io.ReadAll(res.Body) if err != nil { return nil, err } document := &jwksDocument{} if err := json.Unmarshal(body, document); err != nil { return nil, err } keys := make(map[string]interface{}) for _, k := range document.Keys { if k.Use != "" && k.Use != "sig" { continue } switch k.Kty { case "RSA": key, err := rsaKey(k) if err != nil { return nil, err } keys[k.Kid] = key case "EC": key, err := ecKey(k) if err != nil { return nil, err } keys[k.Kid] = key } } return keys, nil } func rsaKey(k jwk) (*rsa.PublicKey, error) { if k.N == "" || k.E == "" { return nil, fmt.Errorf("clé RSA JWKS incomplète (kid %s)", k.Kid) } nBytes, err := base64.RawURLEncoding.DecodeString(k.N) if err != nil { return nil, fmt.Errorf("module RSA invalide: %s", err) } eBytes, err := base64.RawURLEncoding.DecodeString(k.E) if err != nil { return nil, fmt.Errorf("exposant RSA invalide: %s", err) } return &rsa.PublicKey{ N: new(big.Int).SetBytes(nBytes), E: int(new(big.Int).SetBytes(eBytes).Int64()), }, nil } func ecKey(k jwk) (*ecdsa.PublicKey, error) { if k.X == "" || k.Y == "" { return nil, fmt.Errorf("clé EC JWKS incomplète (kid %s)", k.Kid) } var curve elliptic.Curve switch strings.ToUpper(k.Crv) { case "P-256": curve = elliptic.P256() case "P-384": curve = elliptic.P384() case "P-521": curve = elliptic.P521() default: return nil, fmt.Errorf("courbe EC non supportée: %s", k.Crv) } xBytes, err := base64.RawURLEncoding.DecodeString(k.X) if err != nil { return nil, fmt.Errorf("coordonnée x invalide: %s", err) } yBytes, err := base64.RawURLEncoding.DecodeString(k.Y) if err != nil { return nil, fmt.Errorf("coordonnée y invalide: %s", err) } return &ecdsa.PublicKey{ Curve: curve, X: new(big.Int).SetBytes(xBytes), Y: new(big.Int).SetBytes(yBytes), }, nil }