186 lines
4.3 KiB
Go
186 lines
4.3 KiB
Go
package main
|
|
|
|
import (
|
|
"crypto/rsa"
|
|
"encoding/json"
|
|
"errors"
|
|
"fmt"
|
|
"os"
|
|
|
|
"time"
|
|
|
|
"github.com/golang-jwt/jwt/v5"
|
|
)
|
|
|
|
// Load private key from file
|
|
func loadPrivateKey() (*rsa.PrivateKey, error) {
|
|
keyData, err := os.ReadFile("private_key.pem")
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return jwt.ParseRSAPrivateKeyFromPEM(keyData)
|
|
}
|
|
|
|
// Generate a mock JWT with multiple Cognito groups
|
|
func generateMockJWT() (string, error) {
|
|
privateKey, err := loadPrivateKey()
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
|
|
/* payload will look like this when verified
|
|
{
|
|
"cognito:groups": [
|
|
"admins",
|
|
"developers",
|
|
"querybuilders",
|
|
"uploaders",
|
|
"exporters"
|
|
],
|
|
"email": "test@example.com",
|
|
"exp": 1741820939,
|
|
"iat": 1741817339,
|
|
"iss": "https://cognito-idp.us-east-1.amazonaws.com/us-east-1_XXXXXXXXX",
|
|
"sub": "test-user-id"
|
|
}
|
|
*/
|
|
token := jwt.NewWithClaims(jwt.SigningMethodRS256, jwt.MapClaims{
|
|
"sub": "test-user-id",
|
|
"email": "test@example.com",
|
|
"exp": time.Now().Add(time.Hour).Unix(),
|
|
"iat": time.Now().Unix(),
|
|
"iss": "https://cognito-idp.us-east-1.amazonaws.com/us-east-1_XXXXXXXXX",
|
|
"cognito:groups": []string{"admins", "developers", "querybuilders", "uploaders", "exporters"}, // Multiple groups
|
|
})
|
|
|
|
signedToken, errorSigning := token.SignedString(privateKey)
|
|
if errorSigning != nil {
|
|
return "", errorSigning
|
|
}
|
|
fmt.Printf("raw token: %s\n", signedToken)
|
|
|
|
return signedToken, nil
|
|
}
|
|
|
|
// Load public key from file
|
|
func loadPublicKey() (*rsa.PublicKey, error) {
|
|
keyData, err := os.ReadFile("public_key.pem")
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return jwt.ParseRSAPublicKeyFromPEM(keyData)
|
|
}
|
|
|
|
func printTokenDetails(token *jwt.Token) {
|
|
fmt.Println("\n=== JWT Token Details ===")
|
|
|
|
// Print Header
|
|
headerJSON, _ := json.MarshalIndent(token.Header, "", " ")
|
|
fmt.Printf("\nHeader:\n%s\n", string(headerJSON))
|
|
|
|
// Print Claims
|
|
claimsJSON, _ := json.MarshalIndent(token.Claims, "", " ")
|
|
fmt.Printf("\nClaims:\n%s\n", string(claimsJSON))
|
|
|
|
// Print Signature
|
|
fmt.Printf("\nSignature: %s\n", token.Signature)
|
|
|
|
// Print Method
|
|
fmt.Printf("\nSigning Method: %s\n", token.Method.Alg())
|
|
|
|
// Print Valid status
|
|
fmt.Printf("\nValid: %v\n", token.Valid)
|
|
|
|
fmt.Println("\n=====================")
|
|
}
|
|
|
|
// Validate JWT and extract claims
|
|
func validateJWT(tokenString string) (*jwt.Token, jwt.MapClaims, error) {
|
|
publicKey, err := loadPublicKey()
|
|
if err != nil {
|
|
return nil, nil, err
|
|
}
|
|
|
|
token, err := jwt.Parse(tokenString, func(token *jwt.Token) (interface{}, error) {
|
|
if _, ok := token.Method.(*jwt.SigningMethodRSA); !ok {
|
|
return nil, errors.New("unexpected signing method")
|
|
}
|
|
return publicKey, nil
|
|
})
|
|
|
|
// pretty print all of the details of the token
|
|
printTokenDetails(token)
|
|
|
|
if err != nil {
|
|
return nil, nil, err
|
|
}
|
|
|
|
if claims, ok := token.Claims.(jwt.MapClaims); ok && token.Valid {
|
|
return token, claims, nil
|
|
}
|
|
return nil, nil, errors.New("invalid token")
|
|
}
|
|
|
|
func main() {
|
|
tokenString, err := generateMockJWT()
|
|
if err != nil {
|
|
fmt.Println("Error generating token:", err)
|
|
return
|
|
}
|
|
|
|
token, claims, err := validateJWT(tokenString)
|
|
if err != nil {
|
|
fmt.Println("Token validation failed:", err)
|
|
return
|
|
}
|
|
|
|
if token.Valid {
|
|
fmt.Println("Token is valid.")
|
|
fmt.Println("User ID:", claims["sub"])
|
|
fmt.Println("Email:", claims["email"])
|
|
|
|
if groups, ok := claims["cognito:groups"].([]interface{}); ok {
|
|
fmt.Println("User Groups:")
|
|
for _, group := range groups {
|
|
fmt.Println("-", group)
|
|
}
|
|
} else {
|
|
fmt.Println("No groups found in token.")
|
|
}
|
|
} else {
|
|
fmt.Println("Invalid token")
|
|
}
|
|
}
|
|
|
|
//4. Unit Test to Verify Group Membership
|
|
//
|
|
//Add a unit test to verify if the JWT contains expected groups.
|
|
//
|
|
//func TestJWTGroups(t *testing.T) {
|
|
// tokenString, _ := generateMockJWT()
|
|
//
|
|
// _, claims, err := validateJWT(tokenString)
|
|
// if err != nil {
|
|
// t.Errorf("Token validation failed: %v", err)
|
|
// }
|
|
//
|
|
// expectedGroups := []string{"Admins", "Developers", "BetaTesters"}
|
|
// groups, ok := claims["cognito:groups"].([]interface{})
|
|
// if !ok {
|
|
// t.Errorf("Expected 'cognito:groups' claim, but not found")
|
|
// }
|
|
//
|
|
// for _, eg := range expectedGroups {
|
|
// found := false
|
|
// for _, g := range groups {
|
|
// if g.(string) == eg {
|
|
// found = true
|
|
// break
|
|
// }
|
|
// }
|
|
// if !found {
|
|
// t.Errorf("Expected group %s not found in token", eg)
|
|
// }
|
|
// }
|
|
//}
|