package auth import ( "context" "crypto/rand" "fmt" "sync" "time" "github.com/go-chi/jwtauth/v5" ) var TokenAuth *jwtauth.JWTAuth // AuthResponse is the response structure for login/refresh endpoints type AuthResponse struct { Token string `json:"token"` JTI string `json:"jti"` } // In-memory revoked JTI tracking var ( revokedJTIs = make(map[string]time.Time) revokedJTIsMu sync.RWMutex ) // generateJTI generates a UUID v4 string using crypto/rand // Format: xxxxxxxx-xxxx-4xxx-yxxx-xxxxxxxxxxxx func generateJTI() (string, error) { b := make([]byte, 16) if _, err := rand.Read(b); err != nil { return "", fmt.Errorf("failed to generate JTI: %w", err) } // Set version 4 bits b[6] = (b[6] & 0x0f) | 0x40 // Set variant bits (10xx) b[8] = (b[8] & 0x3f) | 0x80 return fmt.Sprintf("%08x-%04x-%04x-%04x-%012x", b[0:4], b[4:6], b[6:8], b[8:10], b[10:16]), nil } // RevokeJTI adds a JTI to the revoked set with its expiry time func RevokeJTI(jti string, expiresAt time.Time) { revokedJTIsMu.Lock() defer revokedJTIsMu.Unlock() revokedJTIs[jti] = expiresAt } // IsJTIRevoked checks if a JTI is in the revoked set func IsJTIRevoked(jti string) bool { revokedJTIsMu.RLock() defer revokedJTIsMu.RUnlock() _, revoked := revokedJTIs[jti] return revoked } // CleanupRevokedJTIs removes entries where the expiry time has passed func CleanupRevokedJTIs() { revokedJTIsMu.Lock() defer revokedJTIsMu.Unlock() now := time.Now() for jti, expiresAt := range revokedJTIs { if now.After(expiresAt) { delete(revokedJTIs, jti) } } } func init() { go func() { ticker := time.NewTicker(5 * time.Minute) defer ticker.Stop() for range ticker.C { CleanupRevokedJTIs() } }() } func InitJWT(secret string) { TokenAuth = jwtauth.New("HS256", []byte(secret), nil) } // GenerateToken creates a JWT with user_id, role, and a unique jti claim // Returns the token string, the JTI, and any error func GenerateToken(userID string, role string) (string, string, error) { jti, err := generateJTI() if err != nil { return "", "", err } _, tokenString, err := TokenAuth.Encode(map[string]interface{}{ "user_id": userID, "role": role, "jti": jti, "exp": time.Now().Add(30 * 24 * time.Hour).Unix(), // 30 days }) return tokenString, jti, err } // VerifyToken validates JWT and returns user_id, role, and jti func VerifyToken(tokenString string, ctx context.Context) (userID string, role string, jti string, err error) { token, err := TokenAuth.Decode(tokenString) if err != nil { return "", "", "", err } var uidVal interface{} if err := token.Get("user_id", &uidVal); err != nil { return "", "", "", fmt.Errorf("invalid user_id claim") } userID, ok := uidVal.(string) if !ok { return "", "", "", fmt.Errorf("invalid user_id claim") } var roleVal interface{} if err := token.Get("role", &roleVal); err != nil { return "", "", "", fmt.Errorf("invalid role claim") } role, ok = roleVal.(string) if !ok { return "", "", "", fmt.Errorf("invalid role claim") } var jtiVal interface{} if err := token.Get("jti", &jtiVal); err != nil { return "", "", "", fmt.Errorf("invalid jti claim") } jti, ok = jtiVal.(string) if !ok || jti == "" { return "", "", "", fmt.Errorf("invalid jti claim") } if IsJTIRevoked(jti) { return "", "", "", fmt.Errorf("token revoked") } jti, ok = jtiVal.(string) if !ok || jti == "" { return "", "", "", fmt.Errorf("invalid jti claim") } if IsJTIRevoked(jti) { return "", "", "", fmt.Errorf("token revoked") } return userID, role, jti, nil }