4
0

restapi_auth_oauth2_test.go 4.4 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159
  1. package otoauth2
  2. import (
  3. "net/http"
  4. "net/http/httptest"
  5. "strconv"
  6. "testing"
  7. "time"
  8. config "github.com/OliveTin/OliveTin/internal/config"
  9. "github.com/stretchr/testify/assert"
  10. "golang.org/x/oauth2"
  11. )
  12. func TestSweepExpiredOAuthStatesLocked(t *testing.T) {
  13. h := &OAuth2Handler{
  14. registeredStates: make(map[string]*oauth2State),
  15. }
  16. h.registeredStates["fresh"] = &oauth2State{
  17. providerName: "test",
  18. createdAt: time.Now(),
  19. }
  20. h.registeredStates["stale"] = &oauth2State{
  21. providerName: "test",
  22. createdAt: time.Now().Add(-2 * oauthStateMaxAge * time.Second),
  23. }
  24. h.sweepExpiredOAuthStatesLocked(time.Now())
  25. _, freshFound := h.registeredStates["fresh"]
  26. _, staleFound := h.registeredStates["stale"]
  27. assert.True(t, freshFound)
  28. assert.False(t, staleFound)
  29. }
  30. func TestGetGroupFieldString(t *testing.T) {
  31. data := map[string]any{"olivetin_group": "admins"}
  32. assert.Equal(t, "admins", getGroupField(data, "olivetin_group", ""))
  33. }
  34. func TestGetGroupFieldMissing(t *testing.T) {
  35. data := map[string]any{}
  36. assert.Equal(t, "", getGroupField(data, "olivetin_group", ""))
  37. }
  38. func TestGetGroupFieldEmptyFieldName(t *testing.T) {
  39. data := map[string]any{"olivetin_group": "admins"}
  40. assert.Equal(t, "", getGroupField(data, "", ""))
  41. }
  42. func TestGetGroupFieldArrayDefaultSeparator(t *testing.T) {
  43. data := map[string]any{"groups": []any{"admins", "ops"}}
  44. assert.Equal(t, "admins ops", getGroupField(data, "groups", ""))
  45. }
  46. func TestGetGroupFieldArrayCustomSeparator(t *testing.T) {
  47. data := map[string]any{"groups": []any{"admins", "ops"}}
  48. assert.Equal(t, "admins,ops", getGroupField(data, "groups", ","))
  49. }
  50. func TestGetGroupFieldArraySkipsNonStringElements(t *testing.T) {
  51. data := map[string]any{"groups": []any{"admins", float64(5), "ops"}}
  52. assert.Equal(t, "admins ops", getGroupField(data, "groups", ""))
  53. }
  54. func TestGetGroupFieldArrayAllNonStringElements(t *testing.T) {
  55. data := map[string]any{"groups": []any{float64(1), true}}
  56. assert.Equal(t, "", getGroupField(data, "groups", ""))
  57. }
  58. func TestGetGroupFieldNotStringOrArray(t *testing.T) {
  59. data := map[string]any{"groups": map[string]any{"nested": "value"}}
  60. assert.Equal(t, "", getGroupField(data, "groups", ""))
  61. }
  62. func TestGetUserInfoWithArrayGroupsClaim(t *testing.T) {
  63. srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
  64. w.Header().Set("Content-Type", "application/json")
  65. _, _ = w.Write([]byte(`{"preferred_username":"john","groups":["admins","ops"]}`))
  66. }))
  67. defer srv.Close()
  68. cfg := config.DefaultConfig()
  69. cfg.AuthHttpHeaderUserGroupSep = ","
  70. provider := &config.OAuth2Provider{
  71. WhoamiUrl: srv.URL,
  72. UsernameField: "preferred_username",
  73. UserGroupField: "groups",
  74. }
  75. userinfo := getUserInfo(cfg, srv.Client(), provider)
  76. assert.Equal(t, "john", userinfo.Username)
  77. assert.Equal(t, "admins,ops", userinfo.Usergroup)
  78. }
  79. func TestComputeUsergroupUsesConfiguredSeparatorWithAddToUsergroup(t *testing.T) {
  80. cfg := config.DefaultConfig()
  81. cfg.AuthHttpHeaderUserGroupSep = ","
  82. h := &OAuth2Handler{cfg: cfg}
  83. userinfo := &UserInfo{Usergroup: "admins,ops"}
  84. providerConfig := &config.OAuth2Provider{AddToUsergroup: "github"}
  85. assert.Equal(t, "admins,ops,github", h.computeUsergroup(userinfo, providerConfig))
  86. }
  87. func TestComputeUsergroupDefaultSeparatorWithAddToUsergroup(t *testing.T) {
  88. cfg := config.DefaultConfig()
  89. h := &OAuth2Handler{cfg: cfg}
  90. userinfo := &UserInfo{Usergroup: "admins"}
  91. providerConfig := &config.OAuth2Provider{AddToUsergroup: "github"}
  92. assert.Equal(t, "admins github", h.computeUsergroup(userinfo, providerConfig))
  93. }
  94. func TestHandleOAuthLoginRejectsWhenStateMapFull(t *testing.T) {
  95. cfg := config.DefaultConfig()
  96. cfg.AuthOAuth2Providers = map[string]*config.OAuth2Provider{
  97. "test": {
  98. Name: "test",
  99. ClientID: "id",
  100. ClientSecret: "secret",
  101. AuthUrl: "https://example.com/auth",
  102. TokenUrl: "https://example.com/token",
  103. },
  104. }
  105. h := NewOAuth2Handler(cfg)
  106. h.registeredStates = make(map[string]*oauth2State, oauthStateMaxEntries)
  107. for i := 0; i < oauthStateMaxEntries; i++ {
  108. h.registeredStates[strconv.Itoa(i)] = &oauth2State{
  109. providerConfig: &oauth2.Config{},
  110. providerName: "test",
  111. createdAt: time.Now(),
  112. }
  113. }
  114. req := httptest.NewRequestWithContext(t.Context(), http.MethodGet, "/oauth/login?provider=test", nil)
  115. rec := httptest.NewRecorder()
  116. h.HandleOAuthLogin(rec, req)
  117. assert.Equal(t, http.StatusServiceUnavailable, rec.Code)
  118. assert.Equal(t, oauthStateMaxEntries, len(h.registeredStates))
  119. }