package usertoken import ( "context" "encoding/json" "errors" "fmt" "strings" "time" "github.com/coreos/go-oidc/v3/oidc" ) type oidcVerifier struct { verifier *oidc.IDTokenVerifier groupsClaim string skew time.Duration } func newOIDCVerifier(ctx context.Context, issuer, audience, groupsClaim string, skew time.Duration) (tokenVerifier, error) { issuer = strings.TrimSpace(issuer) if issuer == "" { return nil, fmt.Errorf("OIDC issuer is empty") } if skew <= 0 { skew = defaultSkew } groupsClaim = strings.TrimSpace(groupsClaim) if groupsClaim == "" { groupsClaim = "groups" } provider, err := oidc.NewProvider(ctx, issuer) if err != nil { return nil, fmt.Errorf("oidc discovery: %w", err) } aud := strings.TrimSpace(audience) verifier := provider.Verifier(&oidc.Config{ ClientID: aud, SkipClientIDCheck: aud == "", SupportedSigningAlgs: []string{"RS256"}, Now: func() time.Time { return time.Now().Add(-skew) }, }) return &oidcVerifier{verifier: verifier, groupsClaim: groupsClaim, skew: skew}, nil } func (v *oidcVerifier) verify(ctx context.Context, raw string) (Caller, error) { if strings.TrimSpace(raw) == "" { return Caller{}, unauthorized("empty token") } tok, err := v.verifier.Verify(ctx, raw) if err != nil { return Caller{}, fmt.Errorf("oidc verify: %w", errors.Join(err, ErrUnauthorized)) } var claims map[string]any if err := tok.Claims(&claims); err != nil { return Caller{}, unauthorized("claims") } now := time.Now() if err := checkTimeClaim(claims, "iat", now, v.skew); err != nil { return Caller{}, err } if err := checkTimeClaim(claims, "nbf", now, v.skew); err != nil { return Caller{}, err } sub, _ := claims["sub"].(string) username, _ := claims["preferred_username"].(string) groups, err := groupsFromMap(claims, v.groupsClaim) if err != nil { return Caller{}, err } return callerFromClaims(sub, username, groups) } // checkTimeClaim rejects a numeric claim that is after now+skew. // A missing claim is fine. func checkTimeClaim(claims map[string]any, name string, now time.Time, skew time.Duration) error { raw, ok := claims[name] if !ok || raw == nil { return nil } sec, ok := numericUnix(raw) if !ok { return unauthorized("bad " + name) } when := time.Unix(sec, 0) if when.After(now.Add(skew)) { return unauthorized(name + " in the future") } return nil } func numericUnix(v any) (int64, bool) { switch n := v.(type) { case float64: return int64(n), true case json.Number: i, err := n.Int64() return i, err == nil default: return 0, false } } func groupsFromMap(claims map[string]any, name string) (*[]string, error) { raw, ok := claims[name] if !ok || raw == nil { return nil, nil } arr, ok := raw.([]any) if !ok { return nil, unauthorized("groups claim") } out := make([]string, 0, len(arr)) for _, item := range arr { s, ok := item.(string) if !ok { return nil, unauthorized("groups claim") } out = append(out, s) } return &out, nil }