4
0

restapi_auth_oauth2.go 13 KB

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