package usertoken import ( "context" "crypto/rand" "crypto/rsa" "crypto/x509" "encoding/base64" "encoding/pem" "errors" "net/http" "net/http/httptest" "os" "path/filepath" "strings" "testing" "time" "github.com/golang-jwt/jwt/v5" ) func testKey(t *testing.T) *rsa.PrivateKey { t.Helper() key, err := rsa.GenerateKey(rand.Reader, 2048) if err != nil { t.Fatalf("generate key: %v", err) } return key } func testAuth(t *testing.T, key *rsa.PrivateKey) *Auth { t.Helper() auth, err := New(t.Context(), Config{ PublicKey: &key.PublicKey, Issuer: "manager", Audience: "mgmnt", }) if err != nil { t.Fatalf("New: %v", err) } return auth } func testSigner(t *testing.T, key *rsa.PrivateKey) *Signer { t.Helper() s, err := NewSigner(SignConfig{PrivateKey: key, Issuer: "manager", Audience: "mgmnt"}) if err != nil { t.Fatalf("NewSigner: %v", err) } return s } func TestMintAndVerify(t *testing.T) { key := testKey(t) signer := testSigner(t, key) auth := testAuth(t, key) groups := []string{"family"} raw, err := signer.Mint(time.Now(), MintInput{Subject: "konrad", Groups: &groups}) if err != nil { t.Fatalf("Mint: %v", err) } caller, err := auth.Verify(t.Context(), raw) if err != nil { t.Fatalf("Verify: %v", err) } if caller.Subject != "konrad" || caller.Username != "konrad" { t.Fatalf("caller = %+v", caller) } if len(caller.Groups) != 1 || caller.Groups[0] != "family" { t.Fatalf("groups = %#v", caller.Groups) } } func TestGroupsOmitVsEmpty(t *testing.T) { key := testKey(t) signer := testSigner(t, key) auth := testAuth(t, key) now := time.Now() omitted, err := signer.Mint(now, MintInput{Subject: "konrad"}) if err != nil { t.Fatalf("Mint omit: %v", err) } c, err := auth.Verify(t.Context(), omitted) if err != nil { t.Fatalf("Verify omit: %v", err) } if c.Groups != nil { t.Fatalf("omitted groups = %#v, want nil", c.Groups) } empty := []string{} present, err := signer.Mint(now, MintInput{Subject: "konrad", Groups: &empty}) if err != nil { t.Fatalf("Mint empty: %v", err) } payload := decodePayload(t, present) if !strings.Contains(payload, `"groups":[]`) { t.Fatalf("payload = %s", payload) } c, err = auth.Verify(t.Context(), present) if err != nil { t.Fatalf("Verify empty: %v", err) } if c.Groups == nil || len(c.Groups) != 0 { t.Fatalf("empty groups = %#v, want empty slice", c.Groups) } } func TestVerifyRejects(t *testing.T) { key := testKey(t) other := testKey(t) signer := testSigner(t, key) auth := testAuth(t, key) now := time.Now() good, err := signer.Mint(now, MintInput{Subject: "konrad"}) if err != nil { t.Fatalf("Mint: %v", err) } otherAuth := testAuth(t, other) if _, err := otherAuth.Verify(t.Context(), good); !errors.Is(err, ErrUnauthorized) { t.Fatalf("wrong key err = %v", err) } expired := signClaims(t, key, tokenClaims{ Issuer: "manager", Subject: "konrad", Audience: jwt.ClaimStrings{"mgmnt"}, IssuedAt: jwt.NewNumericDate(now.Add(-time.Hour)), ExpiresAt: jwt.NewNumericDate(now.Add(-30 * time.Minute)), PreferredUsername: "konrad", }) if _, err := auth.Verify(t.Context(), expired); !errors.Is(err, ErrUnauthorized) { t.Fatalf("expired err = %v", err) } future := signClaims(t, key, tokenClaims{ Issuer: "manager", Subject: "konrad", Audience: jwt.ClaimStrings{"mgmnt"}, IssuedAt: jwt.NewNumericDate(now.Add(time.Hour)), ExpiresAt: jwt.NewNumericDate(now.Add(2 * time.Hour)), PreferredUsername: "konrad", }) if _, err := auth.Verify(t.Context(), future); !errors.Is(err, ErrUnauthorized) { t.Fatalf("future iat err = %v", err) } parts := splitToken(t, good) none := base64.RawURLEncoding.EncodeToString([]byte(`{"alg":"none"}`)) + "." + parts[1] + "." if _, err := auth.Verify(t.Context(), none); !errors.Is(err, ErrUnauthorized) { t.Fatalf("alg none err = %v", err) } } func TestMiddlewareForwardsBearer(t *testing.T) { key := testKey(t) signer := testSigner(t, key) auth := testAuth(t, key) raw, err := signer.Mint(time.Now(), MintInput{Subject: "konrad"}) if err != nil { t.Fatalf("Mint: %v", err) } var gotRaw string var gotSubject string h := auth.Middleware()(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { c, ok := CallerFromContext(r.Context()) if !ok { t.Fatal("missing caller") } gotSubject = c.Subject gotRaw, ok = BearerFromContext(r.Context()) if !ok { t.Fatal("missing bearer") } w.WriteHeader(http.StatusNoContent) })) req := httptest.NewRequest(http.MethodGet, "/v1/lists", nil) req.Header.Set("Authorization", "Bearer "+raw) rec := httptest.NewRecorder() h.ServeHTTP(rec, req) if rec.Code != http.StatusNoContent { t.Fatalf("status = %d", rec.Code) } if gotSubject != "konrad" || gotRaw != raw { t.Fatalf("subject %q raw match %v", gotSubject, gotRaw == raw) } health := httptest.NewRequest(http.MethodGet, "/health", nil) hrec := httptest.NewRecorder() auth.Middleware()(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if _, ok := CallerFromContext(r.Context()); ok { t.Fatal("health has a caller") } w.WriteHeader(http.StatusOK) })).ServeHTTP(hrec, health) if hrec.Code != http.StatusOK { t.Fatalf("health status = %d", hrec.Code) } bad := httptest.NewRequest(http.MethodGet, "/v1/lists", nil) brec := httptest.NewRecorder() h.ServeHTTP(brec, bad) if brec.Code != http.StatusUnauthorized { t.Fatalf("missing bearer status = %d", brec.Code) } } func TestFromEnvModes(t *testing.T) { key := testKey(t) dir := t.TempDir() pubPath := filepath.Join(dir, "user-token.pub") privPath := filepath.Join(dir, "user-token.key") pubDER, err := x509.MarshalPKIXPublicKey(&key.PublicKey) if err != nil { t.Fatal(err) } writePEMBytes(t, pubPath, "PUBLIC KEY", pubDER) writePEMBytes(t, privPath, "RSA PRIVATE KEY", x509.MarshalPKCS1PrivateKey(key)) t.Run("neither", func(t *testing.T) { t.Setenv("USER_TOKEN_PUBLIC_KEY_FILE", "") t.Setenv("OIDC_ISSUER", "") if _, err := FromEnv(context.Background()); err == nil { t.Fatal("expected error") } }) t.Run("both", func(t *testing.T) { t.Setenv("USER_TOKEN_PUBLIC_KEY_FILE", pubPath) t.Setenv("OIDC_ISSUER", "http://issuer.example") if _, err := FromEnv(context.Background()); err == nil { t.Fatal("expected error") } }) t.Run("static", func(t *testing.T) { t.Setenv("USER_TOKEN_PUBLIC_KEY_FILE", pubPath) t.Setenv("OIDC_ISSUER", "") t.Setenv("USER_TOKEN_ISSUER", "manager") t.Setenv("USER_TOKEN_AUDIENCE", "mgmnt") auth, err := FromEnv(t.Context()) if err != nil { t.Fatalf("FromEnv: %v", err) } t.Setenv("USER_TOKEN_PRIVATE_KEY_FILE", privPath) signer, err := NewSignerFromEnv() if err != nil { t.Fatalf("NewSignerFromEnv: %v", err) } raw, err := signer.Mint(time.Now(), MintInput{Subject: "konrad"}) if err != nil { t.Fatal(err) } c, err := auth.Verify(t.Context(), raw) if err != nil { t.Fatalf("Verify: %v", err) } if c.Subject != "konrad" { t.Fatalf("subject = %q", c.Subject) } }) } func signClaims(t *testing.T, key *rsa.PrivateKey, claims tokenClaims) string { t.Helper() s, err := jwt.NewWithClaims(jwt.SigningMethodRS256, claims).SignedString(key) if err != nil { t.Fatalf("sign: %v", err) } return s } func decodePayload(t *testing.T, raw string) string { t.Helper() parts := splitToken(t, raw) b, err := base64.RawURLEncoding.DecodeString(parts[1]) if err != nil { t.Fatalf("payload: %v", err) } return string(b) } func splitToken(t *testing.T, raw string) []string { t.Helper() var parts []string start := 0 for i := 0; i <= len(raw); i++ { if i == len(raw) || raw[i] == '.' { parts = append(parts, raw[start:i]) start = i + 1 } } if len(parts) != 3 { t.Fatalf("token parts = %d", len(parts)) } return parts } func writePEMBytes(t *testing.T, path, typ string, der []byte) { t.Helper() f, err := os.Create(path) if err != nil { t.Fatal(err) } defer f.Close() if err := pem.Encode(f, &pem.Block{Type: typ, Bytes: der}); err != nil { t.Fatal(err) } }