restapi_auth_oauth2.go 11 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447
  1. package otoauth2
  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. "sync"
  14. "time"
  15. authTypes "github.com/OliveTin/OliveTin/internal/auth/authpublic"
  16. config "github.com/OliveTin/OliveTin/internal/config"
  17. log "github.com/sirupsen/logrus"
  18. "golang.org/x/oauth2"
  19. )
  20. type OAuth2Handler struct {
  21. cfg *config.Config
  22. mu sync.RWMutex
  23. registeredStates map[string]*oauth2State
  24. registeredProviders map[string]*oauth2.Config
  25. }
  26. func NewOAuth2Handler(cfg *config.Config) *OAuth2Handler {
  27. h := &OAuth2Handler{
  28. cfg: cfg,
  29. }
  30. h.registeredStates = make(map[string]*oauth2State)
  31. h.registeredProviders = make(map[string]*oauth2.Config)
  32. for providerName, providerConfig := range cfg.AuthOAuth2Providers {
  33. completeProviderConfig(providerName, providerConfig)
  34. newConfig := &oauth2.Config{
  35. ClientID: providerConfig.ClientID,
  36. ClientSecret: providerConfig.ClientSecret,
  37. Scopes: providerConfig.Scopes,
  38. Endpoint: oauth2.Endpoint{
  39. AuthURL: providerConfig.AuthUrl,
  40. TokenURL: providerConfig.TokenUrl,
  41. },
  42. RedirectURL: cfg.AuthOAuth2RedirectURL,
  43. }
  44. h.registeredProviders[providerName] = newConfig
  45. log.Debugf("Dumping newly registered provider: %v = %+v", providerName, providerConfig)
  46. }
  47. return h
  48. }
  49. type oauth2State struct {
  50. providerConfig *oauth2.Config
  51. providerName string
  52. Username string
  53. Usergroup string
  54. createdAt time.Time
  55. }
  56. const (
  57. oauthStateMaxAge = 900 // matches olivetin-sid-oauth cookie MaxAge
  58. oauthStateMaxEntries = 10000
  59. )
  60. func assignIfEmpty(target *string, value string) {
  61. if *target == "" {
  62. *target = value
  63. }
  64. }
  65. func completeProviderConfig(providerName string, providerConfig *config.OAuth2Provider) {
  66. dbConfig, ok := oauth2ProviderDatabase[providerName]
  67. if ok {
  68. assignIfEmpty(&providerConfig.Name, dbConfig.Name)
  69. assignIfEmpty(&providerConfig.Title, dbConfig.Title)
  70. assignIfEmpty(&providerConfig.WhoamiUrl, dbConfig.WhoamiUrl)
  71. assignIfEmpty(&providerConfig.TokenUrl, dbConfig.TokenUrl)
  72. assignIfEmpty(&providerConfig.AuthUrl, dbConfig.AuthUrl)
  73. assignIfEmpty(&providerConfig.Icon, dbConfig.Icon)
  74. assignIfEmpty(&providerConfig.UsernameField, dbConfig.UsernameField)
  75. if providerConfig.Scopes == nil {
  76. providerConfig.Scopes = dbConfig.Scopes
  77. }
  78. } else {
  79. log.Warnf("Provider not found in database: %v", providerName)
  80. }
  81. }
  82. func (h *OAuth2Handler) getOAuth2Config(providerName string) (*oauth2.Config, error) {
  83. config, ok := h.registeredProviders[providerName]
  84. if !ok {
  85. return nil, fmt.Errorf("provider not found in config: %v", providerName)
  86. }
  87. return config, nil
  88. }
  89. func randString(nByte int) (string, error) {
  90. b := make([]byte, nByte)
  91. if _, err := io.ReadFull(rand.Reader, b); err != nil {
  92. return "", err
  93. }
  94. return base64.URLEncoding.EncodeToString(b), nil
  95. }
  96. func (h *OAuth2Handler) cookieSecure(r *http.Request) bool {
  97. useTLS := r.TLS != nil || r.Header.Get("X-Forwarded-Proto") == "https"
  98. return useTLS || h.cfg.Security.ForceSecureCookies
  99. }
  100. func (h *OAuth2Handler) setOAuthCallbackCookie(w http.ResponseWriter, r *http.Request, name, value string) {
  101. cookie := &http.Cookie{
  102. Name: name,
  103. Value: value,
  104. MaxAge: 900, // 15 minutes
  105. Secure: h.cookieSecure(r),
  106. HttpOnly: true,
  107. Path: "/",
  108. SameSite: http.SameSiteLaxMode,
  109. }
  110. http.SetCookie(w, cookie)
  111. }
  112. func (h *OAuth2Handler) deleteOAuthStateLocked(state string) {
  113. delete(h.registeredStates, state)
  114. }
  115. func (h *OAuth2Handler) sweepExpiredOAuthStatesLocked(now time.Time) {
  116. cutoff := now.Add(-oauthStateMaxAge * time.Second)
  117. for state, entry := range h.registeredStates {
  118. if entry.createdAt.Before(cutoff) {
  119. delete(h.registeredStates, state)
  120. }
  121. }
  122. }
  123. func (h *OAuth2Handler) HandleOAuthLogin(w http.ResponseWriter, r *http.Request) {
  124. state, err := randString(16)
  125. if err != nil {
  126. http.Error(w, err.Error(), http.StatusInternalServerError)
  127. return
  128. }
  129. providerName := r.URL.Query().Get("provider")
  130. provider, err := h.getOAuth2Config(providerName)
  131. if err != nil {
  132. log.Errorf("Failed to get provider config: %v %v", providerName, err)
  133. http.Error(w, err.Error(), http.StatusBadRequest)
  134. return
  135. }
  136. h.mu.Lock()
  137. h.sweepExpiredOAuthStatesLocked(time.Now())
  138. if len(h.registeredStates) >= oauthStateMaxEntries {
  139. h.mu.Unlock()
  140. http.Error(w, "OAuth login temporarily unavailable", http.StatusServiceUnavailable)
  141. return
  142. }
  143. h.registeredStates[state] = &oauth2State{
  144. providerConfig: provider,
  145. providerName: providerName,
  146. Username: "",
  147. createdAt: time.Now(),
  148. }
  149. h.mu.Unlock()
  150. h.setOAuthCallbackCookie(w, r, "olivetin-sid-oauth", state)
  151. log.Infof("OAuth2 state: %v mapped to provider %v (found: %v), now redirecting", state, providerName, provider != nil)
  152. http.Redirect(w, r, provider.AuthCodeURL(state), http.StatusFound)
  153. }
  154. func (h *OAuth2Handler) validateStateMatch(queryState, cookieState string) bool {
  155. return queryState == cookieState
  156. }
  157. func (h *OAuth2Handler) checkOAuthCallbackCookie(w http.ResponseWriter, r *http.Request) (*oauth2State, string, bool) {
  158. cookie, err := r.Cookie("olivetin-sid-oauth")
  159. if err != nil {
  160. log.Errorf("Failed to get state cookie: %v", err)
  161. http.Error(w, "State not found", http.StatusBadRequest)
  162. return nil, "", false
  163. }
  164. state := cookie.Value
  165. if !h.validateStateMatch(r.URL.Query().Get("state"), state) {
  166. log.Errorf("State mismatch: %v != %v", r.URL.Query().Get("state"), state)
  167. h.mu.Lock()
  168. h.deleteOAuthStateLocked(state)
  169. h.mu.Unlock()
  170. http.Error(w, "State mismatch", http.StatusBadRequest)
  171. return nil, state, false
  172. }
  173. h.mu.RLock()
  174. registeredState, ok := h.registeredStates[state]
  175. h.mu.RUnlock()
  176. if !ok {
  177. log.Errorf("State not found in server: %v", state)
  178. h.mu.Lock()
  179. h.deleteOAuthStateLocked(state)
  180. h.mu.Unlock()
  181. http.Error(w, "State not found in server", http.StatusBadRequest)
  182. return nil, state, false
  183. }
  184. return registeredState, state, true
  185. }
  186. type HttpClientSettings struct {
  187. Transport *http.Transport
  188. Timeout time.Duration
  189. }
  190. func getOAuth2HttpClient(providerConfig *config.OAuth2Provider) *HttpClientSettings {
  191. config := &HttpClientSettings{
  192. Transport: &http.Transport{
  193. TLSClientConfig: &tls.Config{InsecureSkipVerify: providerConfig.InsecureSkipVerify},
  194. },
  195. Timeout: time.Duration(min(3, providerConfig.CallbackTimeout)) * time.Second,
  196. }
  197. if providerConfig.CertBundlePath != "" {
  198. config.Transport.TLSClientConfig.RootCAs = getOAuthCertBundle(providerConfig)
  199. }
  200. return config
  201. }
  202. func getOAuthCertBundle(providerConfig *config.OAuth2Provider) *x509.CertPool {
  203. caCert, err := os.ReadFile(providerConfig.CertBundlePath)
  204. if err != nil {
  205. log.Errorf("OAuth2 Cert Bundle - failed to read file: %v", err)
  206. return nil
  207. }
  208. caCertPool := x509.NewCertPool()
  209. if ok := caCertPool.AppendCertsFromPEM(caCert); !ok {
  210. log.Errorf("OAuth2 Cert Bundle - failed to append certificates from PEM")
  211. }
  212. return caCertPool
  213. }
  214. func (h *OAuth2Handler) exchangeOAuthCode(ctx context.Context, providerConfig *oauth2.Config, code string, clientSettings *HttpClientSettings) (*oauth2.Token, error) {
  215. exchangeClient := &http.Client{
  216. Transport: clientSettings.Transport,
  217. Timeout: clientSettings.Timeout,
  218. }
  219. ctx = context.WithValue(ctx, oauth2.HTTPClient, exchangeClient)
  220. return providerConfig.Exchange(ctx, code)
  221. }
  222. func (h *OAuth2Handler) createUserInfoClient(ctx context.Context, providerConfig *oauth2.Config, tok *oauth2.Token, clientSettings *HttpClientSettings) *http.Client {
  223. return &http.Client{
  224. Transport: &oauth2.Transport{
  225. Source: providerConfig.TokenSource(ctx, tok),
  226. Base: clientSettings.Transport,
  227. },
  228. Timeout: clientSettings.Timeout,
  229. }
  230. }
  231. func (h *OAuth2Handler) computeUsergroup(userinfo *UserInfo, providerConfig *config.OAuth2Provider) string {
  232. usergroup := userinfo.Usergroup
  233. if providerConfig != nil && providerConfig.AddToUsergroup != "" {
  234. if usergroup != "" {
  235. usergroup = usergroup + " " + providerConfig.AddToUsergroup
  236. } else {
  237. usergroup = providerConfig.AddToUsergroup
  238. }
  239. }
  240. return usergroup
  241. }
  242. func (h *OAuth2Handler) HandleOAuthCallback(w http.ResponseWriter, r *http.Request) {
  243. log.Infof("OAuth2 Callback received")
  244. registeredState, state, ok := h.checkOAuthCallbackCookie(w, r)
  245. if !ok {
  246. return
  247. }
  248. code := r.FormValue("code")
  249. log.WithFields(log.Fields{
  250. "state": state,
  251. "token-code": code,
  252. }).Debug("OAuth2 Token Code")
  253. providerConfig := h.cfg.AuthOAuth2Providers[registeredState.providerName]
  254. clientSettings := getOAuth2HttpClient(providerConfig)
  255. ctx := context.Background()
  256. tok, err := h.exchangeOAuthCode(ctx, registeredState.providerConfig, code, clientSettings)
  257. if err != nil {
  258. log.Errorf("Failed to exchange code: %v", err)
  259. http.Error(w, "Failed to exchange code", http.StatusBadRequest)
  260. return
  261. }
  262. userInfoClient := h.createUserInfoClient(ctx, registeredState.providerConfig, tok, clientSettings)
  263. userinfo := getUserInfo(h.cfg, userInfoClient, providerConfig)
  264. h.mu.Lock()
  265. h.registeredStates[state].Username = userinfo.Username
  266. h.registeredStates[state].Usergroup = h.computeUsergroup(userinfo, providerConfig)
  267. h.mu.Unlock()
  268. http.Redirect(w, r, "/", http.StatusFound)
  269. }
  270. type UserInfo struct {
  271. Username string
  272. Usergroup string
  273. }
  274. //gocyclo:ignore
  275. func getUserInfo(cfg *config.Config, client *http.Client, provider *config.OAuth2Provider) *UserInfo {
  276. ret := &UserInfo{}
  277. res, err := client.Get(provider.WhoamiUrl)
  278. if err != nil {
  279. log.Errorf("Failed to get user data: %v", err)
  280. return ret
  281. }
  282. defer func() { _ = res.Body.Close() }()
  283. if res.StatusCode != http.StatusOK {
  284. log.Errorf("Failed to get user data: %v", res.StatusCode)
  285. return ret
  286. }
  287. contents, err := io.ReadAll(res.Body)
  288. if err != nil {
  289. log.Errorf("Failed to read user data: %v", err)
  290. return ret
  291. }
  292. var userData map[string]any
  293. if cfg.InsecureAllowDumpOAuth2UserData {
  294. log.Debugf("OAuth2 User Data: %v+", string(contents))
  295. }
  296. err = json.Unmarshal(contents, &userData)
  297. if err != nil {
  298. log.Errorf("Failed to unmarshal user data: %v", err)
  299. return ret
  300. }
  301. ret.Username = getDataField(userData, provider.UsernameField)
  302. ret.Usergroup = getDataField(userData, provider.UserGroupField)
  303. return ret
  304. }
  305. func getDataField(data map[string]any, field string) string {
  306. if field == "" {
  307. return ""
  308. }
  309. val, ok := data[field]
  310. if !ok {
  311. log.Errorf("Failed to get field from user data: %v / %v", data, field)
  312. return ""
  313. }
  314. stringVal, ok := val.(string)
  315. if !ok {
  316. log.Errorf("Field %v is not a string: %v", field, val)
  317. return ""
  318. }
  319. return stringVal
  320. }
  321. func (h *OAuth2Handler) lookupOAuth2UserByState(state string) (*authTypes.AuthenticatedUser, bool) {
  322. h.mu.RLock()
  323. serverState, found := h.registeredStates[state]
  324. if !found {
  325. h.mu.RUnlock()
  326. return nil, false
  327. }
  328. user := &authTypes.AuthenticatedUser{
  329. Username: serverState.Username,
  330. UsergroupLine: serverState.Usergroup,
  331. Provider: "oauth2",
  332. SID: state,
  333. }
  334. h.mu.RUnlock()
  335. return user, true
  336. }
  337. func (h *OAuth2Handler) RevokeSession(sid string) {
  338. h.mu.Lock()
  339. defer h.mu.Unlock()
  340. delete(h.registeredStates, sid)
  341. }
  342. func (h *OAuth2Handler) CheckUserFromOAuth2Cookie(context *authTypes.AuthCheckingContext) *authTypes.AuthenticatedUser {
  343. cookie, err := context.Request.Cookie("olivetin-sid-oauth")
  344. if err != nil || cookie.Value == "" {
  345. return nil
  346. }
  347. user, found := h.lookupOAuth2UserByState(cookie.Value)
  348. if !found {
  349. log.WithFields(log.Fields{
  350. "sid": cookie.Value,
  351. "provider": "oauth2",
  352. }).Warnf("Stale session")
  353. return nil
  354. }
  355. return user
  356. }