oauth2.go 5.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249
  1. package httpservers
  2. import (
  3. "context"
  4. "crypto/rand"
  5. "encoding/base64"
  6. "encoding/json"
  7. "fmt"
  8. config "github.com/OliveTin/OliveTin/internal/config"
  9. log "github.com/sirupsen/logrus"
  10. "golang.org/x/oauth2"
  11. "io"
  12. "net/http"
  13. "time"
  14. )
  15. var (
  16. registeredStates = make(map[string]*oauth2State)
  17. registeredProviders = make(map[string]*oauth2.Config)
  18. )
  19. type oauth2State struct {
  20. provider *oauth2.Config
  21. Username string
  22. Usergroup string
  23. }
  24. func assignIfEmpty(target *string, value string) {
  25. if *target == "" {
  26. *target = value
  27. }
  28. }
  29. func completeProviderConfig(providerName string, providerConfig *config.OAuth2Provider) {
  30. dbConfig, ok := oauth2ProviderDatabase[providerName]
  31. if ok {
  32. assignIfEmpty(&providerConfig.WhoamiUrl, dbConfig.WhoamiUrl)
  33. assignIfEmpty(&providerConfig.TokenUrl, dbConfig.TokenUrl)
  34. assignIfEmpty(&providerConfig.AuthUrl, dbConfig.AuthUrl)
  35. assignIfEmpty(&providerConfig.Icon, dbConfig.Icon)
  36. assignIfEmpty(&providerConfig.UsernameField, dbConfig.UsernameField)
  37. if providerConfig.Scopes == nil {
  38. providerConfig.Scopes = dbConfig.Scopes
  39. }
  40. } else {
  41. log.Warnf("Provider not found in database: %v", providerName)
  42. }
  43. }
  44. func getOAuth2Config(cfg *config.Config, providerName string) (*oauth2.Config, error) {
  45. config, ok := registeredProviders[providerName]
  46. if !ok {
  47. providerConfig, ok := cfg.AuthOAuth2Providers[providerName]
  48. if !ok {
  49. return nil, fmt.Errorf("Provider not found in config: %v", providerName)
  50. }
  51. completeProviderConfig(providerName, providerConfig)
  52. config = &oauth2.Config{
  53. ClientID: providerConfig.ClientID,
  54. ClientSecret: providerConfig.ClientSecret,
  55. Scopes: providerConfig.Scopes,
  56. Endpoint: oauth2.Endpoint{
  57. AuthURL: providerConfig.AuthUrl,
  58. TokenURL: providerConfig.TokenUrl,
  59. },
  60. RedirectURL: "http://localhost:1337/oauth/callback",
  61. }
  62. registeredProviders[providerName] = config
  63. log.Debugf("Dumping newly registered provider: %v = %+v", providerName, providerConfig)
  64. }
  65. return config, nil
  66. }
  67. func randString(nByte int) (string, error) {
  68. b := make([]byte, nByte)
  69. if _, err := io.ReadFull(rand.Reader, b); err != nil {
  70. return "", err
  71. }
  72. return base64.URLEncoding.EncodeToString(b), nil
  73. }
  74. func setOauthCallbackCookie(w http.ResponseWriter, r *http.Request, name, value string) {
  75. cookie := &http.Cookie{
  76. Name: name,
  77. Value: value,
  78. MaxAge: int(time.Hour.Seconds()),
  79. Secure: r.TLS != nil,
  80. HttpOnly: true,
  81. Path: "/",
  82. }
  83. http.SetCookie(w, cookie)
  84. }
  85. func handleOAuthLogin(w http.ResponseWriter, r *http.Request) {
  86. state, err := randString(16)
  87. if err != nil {
  88. http.Error(w, err.Error(), http.StatusInternalServerError)
  89. return
  90. }
  91. providerName := r.URL.Query().Get("provider")
  92. provider, err := getOAuth2Config(cfg, providerName)
  93. registeredStates[state] = &oauth2State{
  94. provider: provider,
  95. }
  96. if err != nil {
  97. log.Errorf("Failed to get provider config: %v %v", providerName, err)
  98. http.Error(w, err.Error(), http.StatusBadRequest)
  99. return
  100. }
  101. setOauthCallbackCookie(w, r, "oauth2state", state)
  102. log.Infof("OAuth2 state: %v mapped to provider %v (found: %v), now redirecting", state, providerName, provider != nil)
  103. http.Redirect(w, r, provider.AuthCodeURL(state), http.StatusFound)
  104. }
  105. func checkOAuthCallbackCookie(w http.ResponseWriter, r *http.Request) (*oauth2State, bool) {
  106. state, err := r.Cookie("oauth2state")
  107. if err != nil {
  108. log.Errorf("Failed to get state cookie: %v", err)
  109. http.Error(w, "State not found", http.StatusBadRequest)
  110. return nil, false
  111. }
  112. if r.URL.Query().Get("state") != state.Value {
  113. log.Errorf("State mismatch: %v != %v", r.URL.Query().Get("state"), state.Value)
  114. http.Error(w, "State mismatch", http.StatusBadRequest)
  115. return nil, false
  116. }
  117. registeredState, ok := registeredStates[state.Value]
  118. if !ok {
  119. log.Errorf("State not found in server: %v", state.Value)
  120. http.Error(w, "State not found in server", http.StatusBadRequest)
  121. }
  122. return registeredState, true
  123. }
  124. func handleOAuthCallback(w http.ResponseWriter, r *http.Request) {
  125. log.Infof("OAuth2 Callback received")
  126. registeredState, ok := checkOAuthCallbackCookie(w, r)
  127. if !ok {
  128. return
  129. }
  130. code := r.FormValue("code")
  131. log.Debugf("OAuth2 Token Code: %v", code)
  132. httpClient := &http.Client{Timeout: 2 * time.Second}
  133. ctx := context.Background()
  134. ctx = context.WithValue(ctx, oauth2.HTTPClient, httpClient)
  135. tok, err := registeredState.provider.Exchange(ctx, code)
  136. if err != nil {
  137. log.Errorf("Failed to exchange code: %v", err)
  138. http.Error(w, "Failed to exchange code", http.StatusBadRequest)
  139. return
  140. }
  141. client := registeredState.provider.Client(ctx, tok)
  142. registeredState.Username = getUsername(client)
  143. loginMessage := fmt.Sprintf("Logged in as %v", registeredState.Username)
  144. log.Infof(loginMessage)
  145. w.Write([]byte(loginMessage))
  146. }
  147. func getUsername(client *http.Client) string {
  148. provider := cfg.AuthOAuth2Providers["github"]
  149. res, err := client.Get(provider.WhoamiUrl)
  150. if res.StatusCode != http.StatusOK {
  151. log.Errorf("Failed to get user data: %v", res.StatusCode)
  152. return ""
  153. }
  154. defer res.Body.Close()
  155. contents, err := io.ReadAll(res.Body)
  156. var userData map[string]interface{}
  157. err = json.Unmarshal([]byte(contents), &userData)
  158. if err != nil {
  159. log.Errorf("Failed to unmarshal user data: %v", err)
  160. return ""
  161. }
  162. username, ok := userData[provider.UsernameField]
  163. if !ok {
  164. log.Errorf("Failed to get username from user data: %v", userData)
  165. return ""
  166. }
  167. return username.(string)
  168. }
  169. func parseOAuth2Cookie(r *http.Request) (string, string) {
  170. cookie, err := r.Cookie("oauth2state")
  171. if err != nil {
  172. log.Warnf("Failed to read OAuth2 cookie: %v", err)
  173. return "", ""
  174. }
  175. serverState, found := registeredStates[cookie.Value]
  176. if !found {
  177. log.Warnf("Failed to find OAuth2 state: %v", cookie.Value)
  178. return "", ""
  179. }
  180. return serverState.Username, serverState.Usergroup
  181. }