4
0

restapi_auth_oauth2.go 6.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271
  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. providerConfig *oauth2.Config
  21. providerName string
  22. Username string
  23. Usergroup string
  24. }
  25. func assignIfEmpty(target *string, value string) {
  26. if *target == "" {
  27. *target = value
  28. }
  29. }
  30. func oauth2Init(cfg *config.Config) {
  31. for providerName, providerConfig := range cfg.AuthOAuth2Providers {
  32. completeProviderConfig(providerName, providerConfig)
  33. newConfig := &oauth2.Config{
  34. ClientID: providerConfig.ClientID,
  35. ClientSecret: providerConfig.ClientSecret,
  36. Scopes: providerConfig.Scopes,
  37. Endpoint: oauth2.Endpoint{
  38. AuthURL: providerConfig.AuthUrl,
  39. TokenURL: providerConfig.TokenUrl,
  40. },
  41. RedirectURL: cfg.AuthOAuth2RedirectURL,
  42. }
  43. registeredProviders[providerName] = newConfig
  44. log.Debugf("Dumping newly registered provider: %v = %+v", providerName, providerConfig)
  45. }
  46. }
  47. func completeProviderConfig(providerName string, providerConfig *config.OAuth2Provider) {
  48. dbConfig, ok := oauth2ProviderDatabase[providerName]
  49. if ok {
  50. assignIfEmpty(&providerConfig.Name, dbConfig.Name)
  51. assignIfEmpty(&providerConfig.Title, dbConfig.Title)
  52. assignIfEmpty(&providerConfig.WhoamiUrl, dbConfig.WhoamiUrl)
  53. assignIfEmpty(&providerConfig.TokenUrl, dbConfig.TokenUrl)
  54. assignIfEmpty(&providerConfig.AuthUrl, dbConfig.AuthUrl)
  55. assignIfEmpty(&providerConfig.Icon, dbConfig.Icon)
  56. assignIfEmpty(&providerConfig.UsernameField, dbConfig.UsernameField)
  57. if providerConfig.Scopes == nil {
  58. providerConfig.Scopes = dbConfig.Scopes
  59. }
  60. } else {
  61. log.Warnf("Provider not found in database: %v", providerName)
  62. }
  63. }
  64. func getOAuth2Config(cfg *config.Config, providerName string) (*oauth2.Config, error) {
  65. config, ok := registeredProviders[providerName]
  66. if !ok {
  67. return nil, fmt.Errorf("Provider not found in config: %v", providerName)
  68. }
  69. return config, nil
  70. }
  71. func randString(nByte int) (string, error) {
  72. b := make([]byte, nByte)
  73. if _, err := io.ReadFull(rand.Reader, b); err != nil {
  74. return "", err
  75. }
  76. return base64.URLEncoding.EncodeToString(b), nil
  77. }
  78. func setOauthCallbackCookie(w http.ResponseWriter, r *http.Request, name, value string) {
  79. cookie := &http.Cookie{
  80. Name: name,
  81. Value: value,
  82. MaxAge: 31556952, // 1 year
  83. Secure: r.TLS != nil,
  84. HttpOnly: true,
  85. Path: "/",
  86. }
  87. http.SetCookie(w, cookie)
  88. }
  89. func handleOAuthLogin(w http.ResponseWriter, r *http.Request) {
  90. state, err := randString(16)
  91. if err != nil {
  92. http.Error(w, err.Error(), http.StatusInternalServerError)
  93. return
  94. }
  95. providerName := r.URL.Query().Get("provider")
  96. provider, err := getOAuth2Config(cfg, providerName)
  97. if err != nil {
  98. log.Errorf("Failed to get provider config: %v %v", providerName, err)
  99. http.Error(w, err.Error(), http.StatusBadRequest)
  100. return
  101. }
  102. registeredStates[state] = &oauth2State{
  103. providerConfig: provider,
  104. providerName: providerName,
  105. Username: "",
  106. }
  107. setOauthCallbackCookie(w, r, "olivetin-sid-oauth", state)
  108. log.Infof("OAuth2 state: %v mapped to provider %v (found: %v), now redirecting", state, providerName, provider != nil)
  109. http.Redirect(w, r, provider.AuthCodeURL(state), http.StatusFound)
  110. }
  111. func checkOAuthCallbackCookie(w http.ResponseWriter, r *http.Request) (*oauth2State, string, bool) {
  112. cookie, err := r.Cookie("olivetin-sid-oauth")
  113. state := cookie.Value
  114. if err != nil {
  115. log.Errorf("Failed to get state cookie: %v", err)
  116. http.Error(w, "State not found", http.StatusBadRequest)
  117. return nil, state, false
  118. }
  119. if r.URL.Query().Get("state") != state {
  120. log.Errorf("State mismatch: %v != %v", r.URL.Query().Get("state"), state)
  121. http.Error(w, "State mismatch", http.StatusBadRequest)
  122. return nil, state, false
  123. }
  124. registeredState, ok := registeredStates[state]
  125. if !ok {
  126. log.Errorf("State not found in server: %v", state)
  127. http.Error(w, "State not found in server", http.StatusBadRequest)
  128. }
  129. return registeredState, state, true
  130. }
  131. func handleOAuthCallback(w http.ResponseWriter, r *http.Request) {
  132. log.Infof("OAuth2 Callback received")
  133. registeredState, state, ok := checkOAuthCallbackCookie(w, r)
  134. if !ok {
  135. return
  136. }
  137. code := r.FormValue("code")
  138. log.WithFields(log.Fields{
  139. "state": state,
  140. "token-code": code,
  141. }).Debug("OAuth2 Token Code")
  142. httpClient := &http.Client{Timeout: 2 * time.Second}
  143. ctx := context.Background()
  144. ctx = context.WithValue(ctx, oauth2.HTTPClient, httpClient)
  145. tok, err := registeredState.providerConfig.Exchange(ctx, code)
  146. if err != nil {
  147. log.Errorf("Failed to exchange code: %v", err)
  148. http.Error(w, "Failed to exchange code", http.StatusBadRequest)
  149. return
  150. }
  151. client := registeredState.providerConfig.Client(ctx, tok)
  152. username := getUsername(client, cfg.AuthOAuth2Providers[registeredState.providerName])
  153. registeredStates[state].Username = username
  154. for k, v := range registeredStates {
  155. log.Debugf("states: %+v %+v", k, v)
  156. }
  157. loginMessage := fmt.Sprintf("OAuth2 login complete for %v", registeredStates[state].Username)
  158. log.WithFields(log.Fields{
  159. "state": state,
  160. }).Infof(loginMessage)
  161. http.Redirect(w, r, "/", http.StatusFound)
  162. w.Write([]byte(loginMessage))
  163. }
  164. func getUsername(client *http.Client, provider *config.OAuth2Provider) string {
  165. res, err := client.Get(provider.WhoamiUrl)
  166. if res.StatusCode != http.StatusOK {
  167. log.Errorf("Failed to get user data: %v", res.StatusCode)
  168. return ""
  169. }
  170. defer res.Body.Close()
  171. contents, err := io.ReadAll(res.Body)
  172. var userData map[string]interface{}
  173. err = json.Unmarshal([]byte(contents), &userData)
  174. if err != nil {
  175. log.Errorf("Failed to unmarshal user data: %v", err)
  176. return ""
  177. }
  178. username, ok := userData[provider.UsernameField]
  179. if !ok {
  180. log.Errorf("Failed to get username from user data: %v", userData)
  181. return ""
  182. }
  183. return username.(string)
  184. }
  185. func parseOAuth2Cookie(r *http.Request) (string, string, string) {
  186. cookie, err := r.Cookie("olivetin-sid-oauth")
  187. if err != nil {
  188. log.Warnf("Failed to read OAuth2 cookie: %v", err)
  189. return "", "", ""
  190. }
  191. if cookie.Value == "" {
  192. return "", "", ""
  193. }
  194. serverState, found := registeredStates[cookie.Value]
  195. if !found {
  196. log.Warnf("Failed to find OAuth2 state: %v", cookie.Value)
  197. return "", "", cookie.Value
  198. }
  199. log.Debugf("Found OAuth2 state: %+v", serverState)
  200. return serverState.Username, serverState.Usergroup, cookie.Value
  201. }