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>
This commit is contained in:
+310
@@ -0,0 +1,310 @@
|
||||
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)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user