2026-07-14 17:40:13 +00:00
|
|
|
package auth
|
|
|
|
|
|
|
|
|
|
import (
|
|
|
|
|
"crypto/rsa"
|
|
|
|
|
"crypto/x509"
|
2026-07-28 23:08:32 +00:00
|
|
|
"encoding/base64"
|
|
|
|
|
"encoding/json"
|
2026-07-14 17:40:13 +00:00
|
|
|
"encoding/pem"
|
|
|
|
|
"fmt"
|
2026-07-28 23:08:32 +00:00
|
|
|
"math/big"
|
2026-07-14 17:52:23 +00:00
|
|
|
"net/http"
|
2026-07-14 17:40:13 +00:00
|
|
|
"os"
|
2026-07-14 17:52:23 +00:00
|
|
|
"strings"
|
2026-07-28 23:08:32 +00:00
|
|
|
"time"
|
2026-07-14 17:40:13 +00:00
|
|
|
|
|
|
|
|
"github.com/golang-jwt/jwt/v5"
|
2026-07-28 23:08:32 +00:00
|
|
|
"github.com/rs/zerolog/log"
|
2026-07-14 17:40:13 +00:00
|
|
|
)
|
|
|
|
|
|
|
|
|
|
// RS256Service validates RS256 JWT tokens using a public key
|
|
|
|
|
// Compatible with ouroboros-identity (port 3208) and aamos-admin-v2
|
|
|
|
|
type RS256Service struct {
|
|
|
|
|
publicKey *rsa.PublicKey
|
|
|
|
|
issuer string
|
|
|
|
|
audience string
|
|
|
|
|
}
|
|
|
|
|
|
2026-07-28 23:08:32 +00:00
|
|
|
// JWKS represents a JSON Web Key Set
|
|
|
|
|
type JWKS struct {
|
|
|
|
|
Keys []JWK `json:"keys"`
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
// JWK represents a JSON Web Key
|
|
|
|
|
type JWK struct {
|
|
|
|
|
Kty string `json:"kty"`
|
|
|
|
|
N string `json:"n"`
|
|
|
|
|
E string `json:"e"`
|
|
|
|
|
Use string `json:"use"`
|
|
|
|
|
Alg string `json:"alg"`
|
|
|
|
|
Kid string `json:"kid"`
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
// NewRS256ServiceFromURL fetches JWKS from URL and creates RS256Service
|
|
|
|
|
func NewRS256ServiceFromURL(jwksURL string) (*RS256Service, error) {
|
|
|
|
|
client := &http.Client{Timeout: 10 * time.Second}
|
|
|
|
|
resp, err := client.Get(jwksURL)
|
|
|
|
|
if err != nil {
|
|
|
|
|
return nil, fmt.Errorf("failed to fetch JWKS: %w", err)
|
|
|
|
|
}
|
|
|
|
|
defer resp.Body.Close()
|
|
|
|
|
|
|
|
|
|
if resp.StatusCode != http.StatusOK {
|
|
|
|
|
return nil, fmt.Errorf("JWKS endpoint returned %d", resp.StatusCode)
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
var jwks JWKS
|
|
|
|
|
if err := json.NewDecoder(resp.Body).Decode(&jwks); err != nil {
|
|
|
|
|
return nil, fmt.Errorf("failed to decode JWKS: %w", err)
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
if len(jwks.Keys) == 0 {
|
|
|
|
|
return nil, fmt.Errorf("no keys in JWKS")
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
// Use first signing key
|
|
|
|
|
key := jwks.Keys[0]
|
|
|
|
|
nBytes, err := base64.RawURLEncoding.DecodeString(key.N)
|
|
|
|
|
if err != nil {
|
|
|
|
|
return nil, fmt.Errorf("failed to decode N: %w", err)
|
|
|
|
|
}
|
|
|
|
|
eBytes, err := base64.RawURLEncoding.DecodeString(key.E)
|
|
|
|
|
if err != nil {
|
|
|
|
|
return nil, fmt.Errorf("failed to decode E: %w", err)
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
pub := &rsa.PublicKey{
|
|
|
|
|
N: new(big.Int).SetBytes(nBytes),
|
|
|
|
|
E: int(new(big.Int).SetBytes(eBytes).Int64()),
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
return &RS256Service{
|
|
|
|
|
publicKey: pub,
|
|
|
|
|
issuer: "prexo-identity",
|
|
|
|
|
audience: "prexo",
|
|
|
|
|
}, nil
|
|
|
|
|
}
|
|
|
|
|
|
2026-07-14 17:40:13 +00:00
|
|
|
// NewRS256Service loads the public key from a PEM file
|
|
|
|
|
func NewRS256Service(publicKeyPath string) (*RS256Service, error) {
|
|
|
|
|
pemData, err := os.ReadFile(publicKeyPath)
|
|
|
|
|
if err != nil {
|
|
|
|
|
return nil, fmt.Errorf("failed to read public key: %w", err)
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
block, _ := pem.Decode(pemData)
|
|
|
|
|
if block == nil {
|
|
|
|
|
return nil, fmt.Errorf("failed to decode PEM block")
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
pub, err := x509.ParsePKIXPublicKey(block.Bytes)
|
|
|
|
|
if err != nil {
|
|
|
|
|
// Try PKCS1 format
|
|
|
|
|
pub, err = x509.ParsePKCS1PublicKey(block.Bytes)
|
|
|
|
|
if err != nil {
|
|
|
|
|
return nil, fmt.Errorf("failed to parse public key: %w", err)
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
rsaPub, ok := pub.(*rsa.PublicKey)
|
|
|
|
|
if !ok {
|
|
|
|
|
return nil, fmt.Errorf("not an RSA public key")
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
return &RS256Service{
|
|
|
|
|
publicKey: rsaPub,
|
|
|
|
|
issuer: "prexo-identity",
|
|
|
|
|
audience: "prexo",
|
|
|
|
|
}, nil
|
|
|
|
|
}
|
|
|
|
|
|
2026-07-14 17:52:23 +00:00
|
|
|
// Middleware returns HTTP middleware that validates Bearer tokens using RS256
|
|
|
|
|
func (s *RS256Service) Middleware() func(http.Handler) http.Handler {
|
|
|
|
|
return func(next http.Handler) http.Handler {
|
|
|
|
|
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
|
|
|
authHeader := r.Header.Get("Authorization")
|
|
|
|
|
if authHeader == "" {
|
|
|
|
|
http.Error(w, `{"error":"missing authorization header"}`, http.StatusUnauthorized)
|
|
|
|
|
return
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
if !strings.HasPrefix(authHeader, "Bearer ") {
|
|
|
|
|
http.Error(w, `{"error":"invalid authorization format"}`, http.StatusUnauthorized)
|
|
|
|
|
return
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
tokenString := strings.TrimPrefix(authHeader, "Bearer ")
|
|
|
|
|
claims, err := s.ValidateToken(tokenString)
|
|
|
|
|
if err != nil {
|
2026-07-28 23:08:32 +00:00
|
|
|
log.Warn().Err(err).Msg("token validation failed")
|
2026-07-14 17:52:23 +00:00
|
|
|
http.Error(w, `{"error":"invalid token"}`, http.StatusUnauthorized)
|
|
|
|
|
return
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
ctx := WithClaims(r.Context(), claims)
|
|
|
|
|
next.ServeHTTP(w, r.WithContext(ctx))
|
|
|
|
|
})
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
2026-07-14 17:40:13 +00:00
|
|
|
// ValidateToken verifies an RS256 JWT token
|
|
|
|
|
func (s *RS256Service) ValidateToken(tokenString string) (*Claims, error) {
|
|
|
|
|
token, err := jwt.Parse(tokenString, func(token *jwt.Token) (interface{}, error) {
|
|
|
|
|
if _, ok := token.Method.(*jwt.SigningMethodRSA); !ok {
|
|
|
|
|
return nil, fmt.Errorf("unexpected signing method: %v", token.Header["alg"])
|
|
|
|
|
}
|
|
|
|
|
return s.publicKey, nil
|
|
|
|
|
})
|
2026-07-28 23:08:32 +00:00
|
|
|
if err != nil {
|
|
|
|
|
return nil, fmt.Errorf("token parse error: %w", err)
|
|
|
|
|
}
|
|
|
|
|
if !token.Valid {
|
|
|
|
|
return nil, fmt.Errorf("token invalid")
|
2026-07-14 17:40:13 +00:00
|
|
|
}
|
|
|
|
|
|
|
|
|
|
mapClaims, ok := token.Claims.(jwt.MapClaims)
|
|
|
|
|
if !ok {
|
|
|
|
|
return nil, fmt.Errorf("invalid claims format")
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
claims := &Claims{
|
|
|
|
|
Sub: getStringClaim(mapClaims, "sub"),
|
|
|
|
|
Iss: getStringClaim(mapClaims, "iss"),
|
|
|
|
|
Aud: getStringClaim(mapClaims, "aud"),
|
|
|
|
|
Exp: getInt64Claim(mapClaims, "exp"),
|
|
|
|
|
Iat: getInt64Claim(mapClaims, "iat"),
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
if email, ok := mapClaims["email"].(string); ok {
|
|
|
|
|
claims.Email = email
|
|
|
|
|
}
|
|
|
|
|
if orgID, ok := mapClaims["org_id"].(string); ok {
|
|
|
|
|
claims.OrgID = orgID
|
|
|
|
|
}
|
|
|
|
|
if roles, ok := mapClaims["roles"].([]interface{}); ok {
|
|
|
|
|
claims.Roles = make([]string, len(roles))
|
|
|
|
|
for i, r := range roles {
|
|
|
|
|
claims.Roles[i] = fmt.Sprint(r)
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
if scopes, ok := mapClaims["scopes"].([]interface{}); ok {
|
|
|
|
|
claims.Scopes = make([]string, len(scopes))
|
|
|
|
|
for i, sc := range scopes {
|
|
|
|
|
claims.Scopes[i] = fmt.Sprint(sc)
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
if err := claims.Valid(); err != nil {
|
|
|
|
|
return nil, err
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
return claims, nil
|
|
|
|
|
}
|