restapi_auth_jwt.go 4.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178
  1. package httpservers
  2. import (
  3. "context"
  4. "crypto/rsa"
  5. "errors"
  6. "fmt"
  7. "github.com/golang-jwt/jwt/v5"
  8. log "github.com/sirupsen/logrus"
  9. "net/http"
  10. "os"
  11. "strings"
  12. "github.com/OliveTin/OliveTin/internal/config"
  13. // "github.com/coreos/go-oidc/v3/oidc"
  14. "github.com/MicahParks/keyfunc/v3"
  15. "time"
  16. )
  17. var (
  18. pubKeyBytes []byte = nil
  19. pubKey *rsa.PublicKey
  20. jwksVerifier keyfunc.Keyfunc
  21. )
  22. func initJwks(cfg *config.Config) {
  23. if jwksVerifier == nil {
  24. var err error
  25. if cfg.AuthJwtCertsURL != "" {
  26. ctx, cancel := context.WithTimeout(context.Background(), 300*time.Second)
  27. jwksVerifier, err = keyfunc.NewDefaultCtx(ctx, []string{
  28. cfg.AuthJwtCertsURL,
  29. })
  30. if err != nil {
  31. log.Errorf("Init JWKS Failure: %v", err)
  32. }
  33. defer cancel()
  34. }
  35. }
  36. }
  37. func readLocalPublicKey(cfg *config.Config) error {
  38. if pubKeyBytes != nil {
  39. return nil // Already read.
  40. }
  41. pubKeyBytes, err := os.ReadFile(cfg.AuthJwtPubKeyPath)
  42. if err != nil {
  43. return fmt.Errorf("couldn't read public key from file %s", cfg.AuthJwtPubKeyPath)
  44. }
  45. // Since the token is RSA (which we validated at the start of this function), the return type of this function actually has to be rsa.PublicKey!
  46. pubKey, err = jwt.ParseRSAPublicKeyFromPEM(pubKeyBytes)
  47. if err != nil {
  48. return fmt.Errorf("error parsing public key object (from %s)", cfg.AuthJwtPubKeyPath)
  49. }
  50. return nil
  51. }
  52. func parseJwtTokenWithRemoteKey(cfg *config.Config, jwtToken string) (*jwt.Token, error) {
  53. initJwks(cfg)
  54. return jwt.Parse(jwtToken, jwksVerifier.Keyfunc, jwt.WithAudience(cfg.AuthJwtAud))
  55. }
  56. func parseJwtTokenWithLocalKey(cfg *config.Config, jwtString string) (*jwt.Token, error) {
  57. err := readLocalPublicKey(cfg)
  58. if err != nil {
  59. return nil, err
  60. }
  61. return jwt.Parse(jwtString, func(token *jwt.Token) (interface{}, error) {
  62. if _, ok := token.Method.(*jwt.SigningMethodRSA); !ok {
  63. return nil, fmt.Errorf("parseJwt expected token algorithm RSA but got: %v", token.Header["alg"])
  64. }
  65. return pubKey, nil
  66. })
  67. }
  68. // Hash-based Message Authentication Code
  69. func parseJwtTokenWithHMAC(cfg *config.Config, jwtString string) (*jwt.Token, error) {
  70. return jwt.Parse(jwtString, func(token *jwt.Token) (interface{}, error) {
  71. if _, ok := token.Method.(*jwt.SigningMethodHMAC); !ok {
  72. return nil, fmt.Errorf("parseJwt expected token algorithm HMAC but got: %v", token.Header["alg"])
  73. }
  74. return []byte(cfg.AuthJwtHmacSecret), nil
  75. })
  76. }
  77. func parseJwtToken(cfg *config.Config, jwtString string) (*jwt.Token, error) {
  78. if cfg.AuthJwtCertsURL != "" {
  79. return parseJwtTokenWithRemoteKey(cfg, jwtString)
  80. }
  81. if cfg.AuthJwtPubKeyPath != "" {
  82. return parseJwtTokenWithLocalKey(cfg, jwtString)
  83. }
  84. return parseJwtTokenWithHMAC(cfg, jwtString)
  85. }
  86. func getClaimsFromJwtToken(cfg *config.Config, jwtString string) (jwt.MapClaims, error) {
  87. token, err := parseJwtToken(cfg, jwtString)
  88. if err != nil {
  89. log.Errorf("jwt parse failure: %v", err)
  90. return nil, errors.New("jwt parse failure")
  91. }
  92. if claims, ok := token.Claims.(jwt.MapClaims); ok && token.Valid {
  93. return claims, nil
  94. } else {
  95. return nil, errors.New("jwt token isn't valid")
  96. }
  97. }
  98. func lookupClaimValueOrDefault(claims jwt.MapClaims, key string, def string) string {
  99. if val, ok := claims[key]; ok {
  100. return fmt.Sprintf("%s", val)
  101. } else {
  102. return def
  103. }
  104. }
  105. func parseJwtCookie(cfg *config.Config, request *http.Request) (string, string) {
  106. cookie, err := request.Cookie(cfg.AuthJwtCookieName)
  107. if err != nil {
  108. log.Debugf("jwt cookie check %v name: %v", err, cfg.AuthJwtCookieName)
  109. return "", ""
  110. }
  111. return parseJwt(cfg, cookie.Value)
  112. }
  113. func parseJwt(cfg *config.Config, token string) (string, string) {
  114. claims, err := getClaimsFromJwtToken(cfg, token)
  115. if err != nil {
  116. log.Warnf("jwt claim error: %+v", err)
  117. return "", ""
  118. }
  119. if cfg.InsecureAllowDumpJwtClaims {
  120. log.Debugf("JWT Claims %+v", claims)
  121. }
  122. username := lookupClaimValueOrDefault(claims, cfg.AuthJwtClaimUsername, "")
  123. usergroup := parseGroupClaim(cfg.AuthJwtClaimUserGroup, claims)
  124. return username, usergroup
  125. }
  126. func parseGroupClaim(groupClaim string, claims jwt.MapClaims) string {
  127. usergroup := ""
  128. if val, ok := claims[groupClaim]; ok {
  129. if array, ok := val.([]interface{}); ok {
  130. groups := make([]string, len(array))
  131. for i, v := range array {
  132. groups[i] = fmt.Sprintf("%s", v)
  133. }
  134. usergroup = strings.Join(groups, " ")
  135. } else {
  136. usergroup = fmt.Sprintf("%s", val)
  137. }
  138. }
  139. return usergroup
  140. }