Files
go-usertoken/auth_test.go
T
konradandCursor f082561cc6 Mint and verify short-lived RS256 user tokens.
ExApps sign after AppAPI auth; Microservices check a static public key or an OIDC issuer and forward the same bearer.

Co-authored-by: Cursor <cursoragent@cursor.com>
2026-09-28 14:03:52 +02:00

311 lines
8.0 KiB
Go

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)
}
}