restapi_auth_oauth2.go 9.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381
  1. package httpservers
  2. import (
  3. "context"
  4. "crypto/rand"
  5. "crypto/tls"
  6. "crypto/x509"
  7. "encoding/base64"
  8. "encoding/json"
  9. "fmt"
  10. "io"
  11. "net/http"
  12. "os"
  13. "time"
  14. config "github.com/OliveTin/OliveTin/internal/config"
  15. log "github.com/sirupsen/logrus"
  16. "golang.org/x/oauth2"
  17. )
  18. type OAuth2Handler struct {
  19. cfg *config.Config
  20. registeredStates map[string]*oauth2State
  21. registeredProviders map[string]*oauth2.Config
  22. }
  23. func NewOAuth2Handler(cfg *config.Config) *OAuth2Handler {
  24. h := &OAuth2Handler{
  25. cfg: cfg,
  26. }
  27. h.registeredStates = make(map[string]*oauth2State)
  28. h.registeredProviders = make(map[string]*oauth2.Config)
  29. for providerName, providerConfig := range cfg.AuthOAuth2Providers {
  30. completeProviderConfig(providerName, providerConfig)
  31. newConfig := &oauth2.Config{
  32. ClientID: providerConfig.ClientID,
  33. ClientSecret: providerConfig.ClientSecret,
  34. Scopes: providerConfig.Scopes,
  35. Endpoint: oauth2.Endpoint{
  36. AuthURL: providerConfig.AuthUrl,
  37. TokenURL: providerConfig.TokenUrl,
  38. },
  39. RedirectURL: cfg.AuthOAuth2RedirectURL,
  40. }
  41. h.registeredProviders[providerName] = newConfig
  42. log.Debugf("Dumping newly registered provider: %v = %+v", providerName, providerConfig)
  43. }
  44. return h
  45. }
  46. type oauth2State struct {
  47. providerConfig *oauth2.Config
  48. providerName string
  49. Username string
  50. Usergroup string
  51. }
  52. func assignIfEmpty(target *string, value string) {
  53. if *target == "" {
  54. *target = value
  55. }
  56. }
  57. func completeProviderConfig(providerName string, providerConfig *config.OAuth2Provider) {
  58. dbConfig, ok := oauth2ProviderDatabase[providerName]
  59. if ok {
  60. assignIfEmpty(&providerConfig.Name, dbConfig.Name)
  61. assignIfEmpty(&providerConfig.Title, dbConfig.Title)
  62. assignIfEmpty(&providerConfig.WhoamiUrl, dbConfig.WhoamiUrl)
  63. assignIfEmpty(&providerConfig.TokenUrl, dbConfig.TokenUrl)
  64. assignIfEmpty(&providerConfig.AuthUrl, dbConfig.AuthUrl)
  65. assignIfEmpty(&providerConfig.Icon, dbConfig.Icon)
  66. assignIfEmpty(&providerConfig.UsernameField, dbConfig.UsernameField)
  67. if providerConfig.Scopes == nil {
  68. providerConfig.Scopes = dbConfig.Scopes
  69. }
  70. } else {
  71. log.Warnf("Provider not found in database: %v", providerName)
  72. }
  73. }
  74. func (h *OAuth2Handler) getOAuth2Config(providerName string) (*oauth2.Config, error) {
  75. config, ok := h.registeredProviders[providerName]
  76. if !ok {
  77. return nil, fmt.Errorf("provider not found in config: %v", providerName)
  78. }
  79. return config, nil
  80. }
  81. func randString(nByte int) (string, error) {
  82. b := make([]byte, nByte)
  83. if _, err := io.ReadFull(rand.Reader, b); err != nil {
  84. return "", err
  85. }
  86. return base64.URLEncoding.EncodeToString(b), nil
  87. }
  88. func (h *OAuth2Handler) setOAuthCallbackCookie(w http.ResponseWriter, r *http.Request, name, value string) {
  89. cookie := &http.Cookie{
  90. Name: name,
  91. Value: value,
  92. MaxAge: 31556952, // 1 year
  93. Secure: r.TLS != nil,
  94. HttpOnly: true,
  95. Path: "/",
  96. }
  97. http.SetCookie(w, cookie)
  98. }
  99. func (h *OAuth2Handler) handleOAuthLogin(w http.ResponseWriter, r *http.Request) {
  100. state, err := randString(16)
  101. if err != nil {
  102. http.Error(w, err.Error(), http.StatusInternalServerError)
  103. return
  104. }
  105. providerName := r.URL.Query().Get("provider")
  106. provider, err := h.getOAuth2Config(providerName)
  107. if err != nil {
  108. log.Errorf("Failed to get provider config: %v %v", providerName, err)
  109. http.Error(w, err.Error(), http.StatusBadRequest)
  110. return
  111. }
  112. h.registeredStates[state] = &oauth2State{
  113. providerConfig: provider,
  114. providerName: providerName,
  115. Username: "",
  116. }
  117. h.setOAuthCallbackCookie(w, r, "olivetin-sid-oauth", state)
  118. log.Infof("OAuth2 state: %v mapped to provider %v (found: %v), now redirecting", state, providerName, provider != nil)
  119. http.Redirect(w, r, provider.AuthCodeURL(state), http.StatusFound)
  120. }
  121. func (h *OAuth2Handler) checkOAuthCallbackCookie(w http.ResponseWriter, r *http.Request) (*oauth2State, string, bool) {
  122. cookie, err := r.Cookie("olivetin-sid-oauth")
  123. state := cookie.Value
  124. if err != nil {
  125. log.Errorf("Failed to get state cookie: %v", err)
  126. http.Error(w, "State not found", http.StatusBadRequest)
  127. return nil, state, false
  128. }
  129. if r.URL.Query().Get("state") != state {
  130. log.Errorf("State mismatch: %v != %v", r.URL.Query().Get("state"), state)
  131. http.Error(w, "State mismatch", http.StatusBadRequest)
  132. return nil, state, false
  133. }
  134. registeredState, ok := h.registeredStates[state]
  135. if !ok {
  136. log.Errorf("State not found in server: %v", state)
  137. http.Error(w, "State not found in server", http.StatusBadRequest)
  138. }
  139. return registeredState, state, true
  140. }
  141. type HttpClientSettings struct {
  142. Transport *http.Transport
  143. Timeout time.Duration
  144. }
  145. func getOAuth2HttpClient(providerConfig *config.OAuth2Provider) *HttpClientSettings {
  146. config := &HttpClientSettings{
  147. Transport: &http.Transport{
  148. TLSClientConfig: &tls.Config{InsecureSkipVerify: providerConfig.InsecureSkipVerify},
  149. },
  150. Timeout: time.Duration(min(3, providerConfig.CallbackTimeout)) * time.Second,
  151. }
  152. if providerConfig.CertBundlePath != "" {
  153. config.Transport.TLSClientConfig.RootCAs = getOAuthCertBundle(providerConfig)
  154. }
  155. return config
  156. }
  157. func getOAuthCertBundle(providerConfig *config.OAuth2Provider) *x509.CertPool {
  158. caCert, err := os.ReadFile(providerConfig.CertBundlePath)
  159. if err != nil {
  160. log.Errorf("OAuth2 Cert Bundle - failed to read file: %v", err)
  161. return nil
  162. }
  163. caCertPool := x509.NewCertPool()
  164. if ok := caCertPool.AppendCertsFromPEM(caCert); !ok {
  165. log.Errorf("OAuth2 Cert Bundle - failed to append certificates: %v", err)
  166. }
  167. return caCertPool
  168. }
  169. func (h *OAuth2Handler) handleOAuthCallback(w http.ResponseWriter, r *http.Request) {
  170. log.Infof("OAuth2 Callback received")
  171. registeredState, state, ok := h.checkOAuthCallbackCookie(w, r)
  172. if !ok {
  173. return
  174. }
  175. code := r.FormValue("code")
  176. log.WithFields(log.Fields{
  177. "state": state,
  178. "token-code": code,
  179. }).Debug("OAuth2 Token Code")
  180. providerConfig := h.cfg.AuthOAuth2Providers[registeredState.providerName]
  181. clientSettings := getOAuth2HttpClient(providerConfig)
  182. exchangeClient := &http.Client{
  183. Transport: clientSettings.Transport,
  184. Timeout: clientSettings.Timeout,
  185. }
  186. ctx := context.Background()
  187. ctx = context.WithValue(ctx, oauth2.HTTPClient, exchangeClient)
  188. tok, err := registeredState.providerConfig.Exchange(ctx, code)
  189. if err != nil {
  190. log.Errorf("Failed to exchange code: %v", err)
  191. http.Error(w, "Failed to exchange code", http.StatusBadRequest)
  192. return
  193. }
  194. userInfoClient := &http.Client{
  195. Transport: &oauth2.Transport{
  196. Source: registeredState.providerConfig.TokenSource(ctx, tok),
  197. Base: clientSettings.Transport,
  198. },
  199. Timeout: clientSettings.Timeout,
  200. }
  201. userinfo := getUserInfo(h.cfg, userInfoClient, h.cfg.AuthOAuth2Providers[registeredState.providerName])
  202. h.registeredStates[state].Username = userinfo.Username
  203. h.registeredStates[state].Usergroup = userinfo.Usergroup
  204. for k, v := range h.registeredStates {
  205. log.Debugf("states: %+v %+v", k, v)
  206. }
  207. log.WithFields(log.Fields{
  208. "state": state,
  209. "username": h.registeredStates[state].Username,
  210. }).Info("OAuth2 login successful")
  211. http.Redirect(w, r, "/", http.StatusFound)
  212. w.Write([]byte("OAuth2 login successful."))
  213. }
  214. type UserInfo struct {
  215. Username string
  216. Usergroup string
  217. }
  218. //gocyclo:ignore
  219. func getUserInfo(cfg *config.Config, client *http.Client, provider *config.OAuth2Provider) *UserInfo {
  220. ret := &UserInfo{}
  221. res, err := client.Get(provider.WhoamiUrl)
  222. if err != nil {
  223. log.Errorf("Failed to get user data: %v", err)
  224. return ret
  225. }
  226. if res.StatusCode != http.StatusOK {
  227. log.Errorf("Failed to get user data: %v", res.StatusCode)
  228. return ret
  229. }
  230. defer res.Body.Close()
  231. contents, err := io.ReadAll(res.Body)
  232. if err != nil {
  233. log.Errorf("Failed to read user data: %v", err)
  234. return ret
  235. }
  236. var userData map[string]any
  237. if cfg.InsecureAllowDumpOAuth2UserData {
  238. log.Debugf("OAuth2 User Data: %v+", string(contents))
  239. }
  240. err = json.Unmarshal([]byte(contents), &userData)
  241. if err != nil {
  242. log.Errorf("Failed to unmarshal user data: %v", err)
  243. return ret
  244. }
  245. ret.Username = getDataField(userData, provider.UsernameField)
  246. ret.Usergroup = getDataField(userData, provider.UserGroupField)
  247. return ret
  248. }
  249. func getDataField(data map[string]any, field string) string {
  250. if field == "" {
  251. return ""
  252. }
  253. val, ok := data[field]
  254. if !ok {
  255. log.Errorf("Failed to get field from user data: %v / %v", data, field)
  256. return ""
  257. }
  258. stringVal, ok := val.(string)
  259. if !ok {
  260. log.Errorf("Field %v is not a string: %v", field, val)
  261. return ""
  262. }
  263. return stringVal
  264. }
  265. func (h *OAuth2Handler) parseOAuth2Cookie(r *http.Request) (string, string, string) {
  266. cookie, err := r.Cookie("olivetin-sid-oauth")
  267. if err != nil {
  268. log.Warnf("Failed to read OAuth2 cookie: %v", err)
  269. return "", "", ""
  270. }
  271. if cookie.Value == "" {
  272. return "", "", ""
  273. }
  274. serverState, found := h.registeredStates[cookie.Value]
  275. if !found {
  276. log.WithFields(log.Fields{
  277. "sid": cookie.Value,
  278. "provider": "oauth2",
  279. }).Warnf("Stale session")
  280. return "", "", cookie.Value
  281. }
  282. log.Debugf("Found OAuth2 state: %+v", serverState)
  283. return serverState.Username, serverState.Usergroup, cookie.Value
  284. }