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>
311 lines
8.0 KiB
Go
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)
|
|
}
|
|
}
|