package auth import ( "crypto/rand" "crypto/rsa" "crypto/x509" "encoding/pem" "net/http" "net/http/httptest" "os" "testing" "time" "github.com/golang-jwt/jwt/v5" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) func generateTestKeyPair(t *testing.T) (privateKey *rsa.PrivateKey, publicKeyPEM []byte) { privateKey, err := rsa.GenerateKey(rand.Reader, 2048) require.NoError(t, err) publicKeyBytes, err := x509.MarshalPKIXPublicKey(&privateKey.PublicKey) require.NoError(t, err) publicKeyPEM = pem.EncodeToMemory(&pem.Block{ Type: "PUBLIC KEY", Bytes: publicKeyBytes, }) return privateKey, publicKeyPEM } func TestNewRS256Service(t *testing.T) { _, pubPEM := generateTestKeyPair(t) tmpFile, err := os.CreateTemp("", "test-pub-*.pem") require.NoError(t, err) defer os.Remove(tmpFile.Name()) _, err = tmpFile.Write(pubPEM) require.NoError(t, err) tmpFile.Close() svc, err := NewRS256Service(tmpFile.Name()) require.NoError(t, err) assert.NotNil(t, svc.publicKey) assert.Equal(t, "prexo-identity", svc.issuer) assert.Equal(t, "prexo", svc.audience) } func TestRS256Service_ValidateToken_Success(t *testing.T) { privateKey, pubPEM := generateTestKeyPair(t) tmpFile, err := os.CreateTemp("", "test-pub-*.pem") require.NoError(t, err) defer os.Remove(tmpFile.Name()) _, err = tmpFile.Write(pubPEM) require.NoError(t, err) tmpFile.Close() svc, err := NewRS256Service(tmpFile.Name()) require.NoError(t, err) // Issue a token with the private key now := time.Now().Unix() claims := jwt.MapClaims{ "sub": "user-123", "email": "test@example.com", "org_id": "org-456", "roles": []string{"admin", "viewer"}, "iss": "prexo-identity", "aud": "prexo", "iat": now, "exp": now + 3600, } token := jwt.NewWithClaims(jwt.SigningMethodRS256, claims) tokenString, err := token.SignedString(privateKey) require.NoError(t, err) // Validate with the service validated, err := svc.ValidateToken(tokenString) require.NoError(t, err) assert.Equal(t, "user-123", validated.Sub) assert.Equal(t, "test@example.com", validated.Email) assert.Equal(t, "org-456", validated.OrgID) assert.Equal(t, []string{"admin", "viewer"}, validated.Roles) } func TestRS256Service_ValidateToken_InvalidSignature(t *testing.T) { // Generate two different key pairs _, pubPEM1 := generateTestKeyPair(t) privateKey2, _ := generateTestKeyPair(t) tmpFile, err := os.CreateTemp("", "test-pub-*.pem") require.NoError(t, err) defer os.Remove(tmpFile.Name()) _, err = tmpFile.Write(pubPEM1) require.NoError(t, err) tmpFile.Close() svc, err := NewRS256Service(tmpFile.Name()) require.NoError(t, err) // Sign with key 2, validate with key 1 now := time.Now().Unix() claims := jwt.MapClaims{ "sub": "user-123", "iss": "prexo-identity", "aud": "prexo", "iat": now, "exp": now + 3600, } token := jwt.NewWithClaims(jwt.SigningMethodRS256, claims) tokenString, err := token.SignedString(privateKey2) require.NoError(t, err) _, err = svc.ValidateToken(tokenString) assert.Error(t, err) } func TestRS256Service_ValidateToken_Expired(t *testing.T) { privateKey, pubPEM := generateTestKeyPair(t) tmpFile, err := os.CreateTemp("", "test-pub-*.pem") require.NoError(t, err) defer os.Remove(tmpFile.Name()) _, err = tmpFile.Write(pubPEM) require.NoError(t, err) tmpFile.Close() svc, err := NewRS256Service(tmpFile.Name()) require.NoError(t, err) // Issue expired token claims := jwt.MapClaims{ "sub": "user-123", "iss": "prexo-identity", "aud": "prexo", "iat": time.Now().Unix() - 7200, "exp": time.Now().Unix() - 3600, } token := jwt.NewWithClaims(jwt.SigningMethodRS256, claims) tokenString, err := token.SignedString(privateKey) require.NoError(t, err) _, err = svc.ValidateToken(tokenString) assert.Error(t, err) } func TestRS256Service_Middleware_ValidToken(t *testing.T) { privateKey, pubPEM := generateTestKeyPair(t) tmpFile, err := os.CreateTemp("", "test-pub-*.pem") require.NoError(t, err) defer os.Remove(tmpFile.Name()) _, err = tmpFile.Write(pubPEM) require.NoError(t, err) tmpFile.Close() svc, err := NewRS256Service(tmpFile.Name()) require.NoError(t, err) // Issue a valid token now := time.Now().Unix() claims := jwt.MapClaims{ "sub": "user-123", "email": "test@example.com", "iss": "prexo-identity", "aud": "prexo", "iat": now, "exp": now + 3600, } token := jwt.NewWithClaims(jwt.SigningMethodRS256, claims) tokenString, err := token.SignedString(privateKey) require.NoError(t, err) // Test middleware handler := svc.Middleware()(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { validatedClaims, ok := FromContext(r.Context()) require.True(t, ok) assert.Equal(t, "user-123", validatedClaims.Sub) w.WriteHeader(http.StatusOK) })) req := httptest.NewRequest(http.MethodGet, "/api/test", nil) req.Header.Set("Authorization", "Bearer "+tokenString) rr := httptest.NewRecorder() handler.ServeHTTP(rr, req) assert.Equal(t, http.StatusOK, rr.Code) } func TestRS256Service_Middleware_InvalidToken(t *testing.T) { _, pubPEM := generateTestKeyPair(t) tmpFile, err := os.CreateTemp("", "test-pub-*.pem") require.NoError(t, err) defer os.Remove(tmpFile.Name()) _, err = tmpFile.Write(pubPEM) require.NoError(t, err) tmpFile.Close() svc, err := NewRS256Service(tmpFile.Name()) require.NoError(t, err) handler := svc.Middleware()(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { t.Fatal("should not reach handler") })) req := httptest.NewRequest(http.MethodGet, "/api/test", nil) req.Header.Set("Authorization", "Bearer invalid-token") rr := httptest.NewRecorder() handler.ServeHTTP(rr, req) assert.Equal(t, http.StatusUnauthorized, rr.Code) } func TestRS256Service_ValidateToken_HS256(t *testing.T) { _, pubPEM := generateTestKeyPair(t) tmpFile, err := os.CreateTemp("", "test-pub-*.pem") require.NoError(t, err) defer os.Remove(tmpFile.Name()) _, err = tmpFile.Write(pubPEM) require.NoError(t, err) tmpFile.Close() svc, err := NewRS256Service(tmpFile.Name()) require.NoError(t, err) // Sign with HS256 instead of RS256 claims := jwt.MapClaims{ "sub": "user-123", "iss": "prexo-identity", "aud": "prexo", "iat": time.Now().Unix(), "exp": time.Now().Unix() + 3600, } token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims) tokenString, err := token.SignedString([]byte("secret")) require.NoError(t, err) _, err = svc.ValidateToken(tokenString) assert.Error(t, err) assert.Contains(t, err.Error(), "unexpected signing method") }