oidc_test.go 57 KB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061626364656667686970717273747576777879808182838485868788899091929394959697989910010110210310410510610710810911011111211311411511611711811912012112212312412512612712812913013113213313413513613713813914014114214314414514614714814915015115215315415515615715815916016116216316416516616716816917017117217317417517617717817918018118218318418518618718818919019119219319419519619719819920020120220320420520620720820921021121221321421521621721821922022122222322422522622722822923023123223323423523623723823924024124224324424524624724824925025125225325425525625725825926026126226326426526626726826927027127227327427527627727827928028128228328428528628728828929029129229329429529629729829930030130230330430530630730830931031131231331431531631731831932032132232332432532632732832933033133233333433533633733833934034134234334434534634734834935035135235335435535635735835936036136236336436536636736836937037137237337437537637737837938038138238338438538638738838939039139239339439539639739839940040140240340440540640740840941041141241341441541641741841942042142242342442542642742842943043143243343443543643743843944044144244344444544644744844945045145245345445545645745845946046146246346446546646746846947047147247347447547647747847948048148248348448548648748848949049149249349449549649749849950050150250350450550650750850951051151251351451551651751851952052152252352452552652752852953053153253353453553653753853954054154254354454554654754854955055155255355455555655755855956056156256356456556656756856957057157257357457557657757857958058158258358458558658758858959059159259359459559659759859960060160260360460560660760860961061161261361461561661761861962062162262362462562662762862963063163263363463563663763863964064164264364464564664764864965065165265365465565665765865966066166266366466566666766866967067167267367467567667767867968068168268368468568668768868969069169269369469569669769869970070170270370470570670770870971071171271371471571671771871972072172272372472572672772872973073173273373473573673773873974074174274374474574674774874975075175275375475575675775875976076176276376476576676776876977077177277377477577677777877978078178278378478578678778878979079179279379479579679779879980080180280380480580680780880981081181281381481581681781881982082182282382482582682782882983083183283383483583683783883984084184284384484584684784884985085185285385485585685785885986086186286386486586686786886987087187287387487587687787887988088188288388488588688788888989089189289389489589689789889990090190290390490590690790890991091191291391491591691791891992092192292392492592692792892993093193293393493593693793893994094194294394494594694794894995095195295395495595695795895996096196296396496596696796896997097197297397497597697797897998098198298398498598698798898999099199299399499599699799899910001001100210031004100510061007100810091010101110121013101410151016101710181019102010211022102310241025102610271028102910301031103210331034103510361037103810391040104110421043104410451046104710481049105010511052105310541055105610571058105910601061106210631064106510661067106810691070107110721073107410751076107710781079108010811082108310841085108610871088108910901091109210931094109510961097109810991100110111021103110411051106110711081109111011111112111311141115111611171118111911201121112211231124112511261127112811291130113111321133113411351136113711381139114011411142114311441145114611471148114911501151115211531154115511561157115811591160116111621163116411651166116711681169117011711172117311741175117611771178117911801181118211831184118511861187118811891190119111921193119411951196119711981199120012011202120312041205120612071208120912101211121212131214121512161217121812191220122112221223122412251226122712281229123012311232123312341235123612371238123912401241124212431244124512461247124812491250125112521253125412551256125712581259126012611262126312641265126612671268126912701271127212731274127512761277127812791280128112821283128412851286128712881289129012911292129312941295129612971298129913001301130213031304130513061307130813091310131113121313131413151316131713181319132013211322132313241325132613271328132913301331133213331334133513361337133813391340134113421343134413451346134713481349135013511352135313541355135613571358135913601361136213631364136513661367136813691370137113721373137413751376137713781379138013811382138313841385138613871388138913901391139213931394139513961397139813991400140114021403140414051406140714081409141014111412141314141415141614171418141914201421142214231424142514261427142814291430143114321433143414351436143714381439144014411442144314441445144614471448144914501451145214531454145514561457145814591460146114621463146414651466146714681469147014711472147314741475147614771478147914801481148214831484148514861487148814891490149114921493149414951496149714981499150015011502150315041505150615071508150915101511151215131514151515161517151815191520152115221523152415251526152715281529153015311532153315341535153615371538153915401541154215431544154515461547154815491550155115521553155415551556155715581559156015611562156315641565156615671568156915701571157215731574157515761577157815791580158115821583158415851586158715881589159015911592159315941595159615971598159916001601160216031604160516061607160816091610161116121613161416151616161716181619162016211622162316241625162616271628162916301631163216331634163516361637163816391640164116421643164416451646164716481649165016511652165316541655165616571658165916601661166216631664166516661667166816691670167116721673167416751676167716781679168016811682168316841685168616871688168916901691169216931694169516961697169816991700170117021703170417051706170717081709171017111712171317141715171617171718171917201721172217231724172517261727172817291730173117321733173417351736173717381739174017411742174317441745174617471748174917501751175217531754175517561757175817591760176117621763176417651766176717681769177017711772
  1. // Copyright (C) 2019 Nicola Murino
  2. //
  3. // This program is free software: you can redistribute it and/or modify
  4. // it under the terms of the GNU Affero General Public License as published
  5. // by the Free Software Foundation, version 3.
  6. //
  7. // This program is distributed in the hope that it will be useful,
  8. // but WITHOUT ANY WARRANTY; without even the implied warranty of
  9. // MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
  10. // GNU Affero General Public License for more details.
  11. //
  12. // You should have received a copy of the GNU Affero General Public License
  13. // along with this program. If not, see <https://www.gnu.org/licenses/>.
  14. package httpd
  15. import (
  16. "bytes"
  17. "context"
  18. "encoding/json"
  19. "fmt"
  20. "io/fs"
  21. "net/http"
  22. "net/http/httptest"
  23. "net/url"
  24. "os"
  25. "path/filepath"
  26. "reflect"
  27. "runtime"
  28. "testing"
  29. "time"
  30. "unsafe"
  31. "github.com/coreos/go-oidc/v3/oidc"
  32. "github.com/go-chi/jwtauth/v5"
  33. "github.com/rs/xid"
  34. "github.com/sftpgo/sdk"
  35. "github.com/stretchr/testify/assert"
  36. "github.com/stretchr/testify/require"
  37. "golang.org/x/oauth2"
  38. "github.com/drakkan/sftpgo/v2/internal/common"
  39. "github.com/drakkan/sftpgo/v2/internal/dataprovider"
  40. "github.com/drakkan/sftpgo/v2/internal/kms"
  41. "github.com/drakkan/sftpgo/v2/internal/util"
  42. "github.com/drakkan/sftpgo/v2/internal/vfs"
  43. )
  44. const (
  45. oidcMockAddr = "127.0.0.1:11111"
  46. )
  47. type mockTokenSource struct {
  48. token *oauth2.Token
  49. err error
  50. }
  51. func (t *mockTokenSource) Token() (*oauth2.Token, error) {
  52. return t.token, t.err
  53. }
  54. type mockOAuth2Config struct {
  55. tokenSource *mockTokenSource
  56. authCodeURL string
  57. token *oauth2.Token
  58. err error
  59. }
  60. func (c *mockOAuth2Config) AuthCodeURL(_ string, _ ...oauth2.AuthCodeOption) string {
  61. return c.authCodeURL
  62. }
  63. func (c *mockOAuth2Config) Exchange(_ context.Context, _ string, _ ...oauth2.AuthCodeOption) (*oauth2.Token, error) {
  64. return c.token, c.err
  65. }
  66. func (c *mockOAuth2Config) TokenSource(_ context.Context, _ *oauth2.Token) oauth2.TokenSource {
  67. return c.tokenSource
  68. }
  69. type mockOIDCVerifier struct {
  70. token *oidc.IDToken
  71. err error
  72. }
  73. func (v *mockOIDCVerifier) Verify(_ context.Context, _ string) (*oidc.IDToken, error) {
  74. return v.token, v.err
  75. }
  76. // hack because the field is unexported
  77. func setIDTokenClaims(idToken *oidc.IDToken, claims []byte) {
  78. pointerVal := reflect.ValueOf(idToken)
  79. val := reflect.Indirect(pointerVal)
  80. member := val.FieldByName("claims")
  81. ptr := unsafe.Pointer(member.UnsafeAddr())
  82. realPtr := (*[]byte)(ptr)
  83. *realPtr = claims
  84. }
  85. func TestOIDCInitialization(t *testing.T) {
  86. config := OIDC{}
  87. err := config.initialize()
  88. assert.NoError(t, err)
  89. secret := "jRsmE0SWnuZjP7djBqNq0mrf8QN77j2c"
  90. config = OIDC{
  91. ClientID: "sftpgo-client",
  92. ClientSecret: util.GenerateUniqueID(),
  93. ConfigURL: fmt.Sprintf("http://%v/", oidcMockAddr),
  94. RedirectBaseURL: "http://127.0.0.1:8081/",
  95. UsernameField: "preferred_username",
  96. RoleField: "sftpgo_role",
  97. }
  98. err = config.initialize()
  99. if assert.Error(t, err) {
  100. assert.Contains(t, err.Error(), "oidc: required scope \"openid\" is not set")
  101. }
  102. config.Scopes = []string{oidc.ScopeOpenID}
  103. config.ClientSecretFile = "missing file"
  104. err = config.initialize()
  105. assert.ErrorIs(t, err, fs.ErrNotExist)
  106. secretFile := filepath.Join(os.TempDir(), util.GenerateUniqueID())
  107. defer os.Remove(secretFile)
  108. err = os.WriteFile(secretFile, []byte(secret), 0600)
  109. assert.NoError(t, err)
  110. config.ClientSecretFile = secretFile
  111. err = config.initialize()
  112. if assert.Error(t, err) {
  113. assert.Contains(t, err.Error(), "oidc: unable to initialize provider")
  114. }
  115. assert.Equal(t, secret, config.ClientSecret)
  116. config.ConfigURL = fmt.Sprintf("http://%v/auth/realms/sftpgo", oidcMockAddr)
  117. err = config.initialize()
  118. assert.NoError(t, err)
  119. assert.Equal(t, "http://127.0.0.1:8081"+webOIDCRedirectPath, config.getRedirectURL())
  120. }
  121. func TestOIDCLoginLogout(t *testing.T) {
  122. oidcMgr, ok := oidcMgr.(*memoryOIDCManager)
  123. require.True(t, ok)
  124. server := getTestOIDCServer()
  125. err := server.binding.OIDC.initialize()
  126. assert.NoError(t, err)
  127. server.initializeRouter()
  128. rr := httptest.NewRecorder()
  129. r, err := http.NewRequest(http.MethodGet, webOIDCRedirectPath, nil)
  130. assert.NoError(t, err)
  131. server.router.ServeHTTP(rr, r)
  132. assert.Equal(t, http.StatusBadRequest, rr.Code)
  133. assert.Contains(t, rr.Body.String(), util.I18nInvalidAuth)
  134. expiredAuthReq := oidcPendingAuth{
  135. State: xid.New().String(),
  136. Nonce: xid.New().String(),
  137. Audience: tokenAudienceWebClient,
  138. IssuedAt: util.GetTimeAsMsSinceEpoch(time.Now().Add(-10 * time.Minute)),
  139. }
  140. oidcMgr.addPendingAuth(expiredAuthReq)
  141. rr = httptest.NewRecorder()
  142. r, err = http.NewRequest(http.MethodGet, webOIDCRedirectPath+"?state="+expiredAuthReq.State, nil)
  143. assert.NoError(t, err)
  144. server.router.ServeHTTP(rr, r)
  145. assert.Equal(t, http.StatusBadRequest, rr.Code)
  146. assert.Contains(t, rr.Body.String(), util.I18nInvalidAuth)
  147. oidcMgr.removePendingAuth(expiredAuthReq.State)
  148. server.binding.OIDC.oauth2Config = &mockOAuth2Config{
  149. tokenSource: &mockTokenSource{},
  150. authCodeURL: webOIDCRedirectPath,
  151. err: common.ErrGenericFailure,
  152. }
  153. server.binding.OIDC.verifier = &mockOIDCVerifier{
  154. err: common.ErrGenericFailure,
  155. }
  156. rr = httptest.NewRecorder()
  157. r, err = http.NewRequest(http.MethodGet, webAdminOIDCLoginPath, nil)
  158. assert.NoError(t, err)
  159. server.router.ServeHTTP(rr, r)
  160. assert.Equal(t, http.StatusFound, rr.Code)
  161. assert.Equal(t, webOIDCRedirectPath, rr.Header().Get("Location"))
  162. require.Len(t, oidcMgr.pendingAuths, 1)
  163. var state string
  164. for k := range oidcMgr.pendingAuths {
  165. state = k
  166. }
  167. rr = httptest.NewRecorder()
  168. r, err = http.NewRequest(http.MethodGet, webOIDCRedirectPath+"?state="+state, nil)
  169. assert.NoError(t, err)
  170. server.router.ServeHTTP(rr, r)
  171. assert.Equal(t, http.StatusFound, rr.Code)
  172. assert.Equal(t, webAdminLoginPath, rr.Header().Get("Location"))
  173. require.Len(t, oidcMgr.pendingAuths, 0)
  174. rr = httptest.NewRecorder()
  175. r, err = http.NewRequest(http.MethodGet, webAdminLoginPath, nil)
  176. assert.NoError(t, err)
  177. server.router.ServeHTTP(rr, r)
  178. assert.Equal(t, http.StatusOK, rr.Code)
  179. // now the same for the web client
  180. rr = httptest.NewRecorder()
  181. r, err = http.NewRequest(http.MethodGet, webClientOIDCLoginPath, nil)
  182. assert.NoError(t, err)
  183. server.router.ServeHTTP(rr, r)
  184. assert.Equal(t, http.StatusFound, rr.Code)
  185. assert.Equal(t, webOIDCRedirectPath, rr.Header().Get("Location"))
  186. require.Len(t, oidcMgr.pendingAuths, 1)
  187. for k := range oidcMgr.pendingAuths {
  188. state = k
  189. }
  190. rr = httptest.NewRecorder()
  191. r, err = http.NewRequest(http.MethodGet, webOIDCRedirectPath+"?state="+state, nil)
  192. assert.NoError(t, err)
  193. server.router.ServeHTTP(rr, r)
  194. assert.Equal(t, http.StatusFound, rr.Code)
  195. assert.Equal(t, webClientLoginPath, rr.Header().Get("Location"))
  196. require.Len(t, oidcMgr.pendingAuths, 0)
  197. rr = httptest.NewRecorder()
  198. r, err = http.NewRequest(http.MethodGet, webClientLoginPath, nil)
  199. assert.NoError(t, err)
  200. server.router.ServeHTTP(rr, r)
  201. assert.Equal(t, http.StatusOK, rr.Code)
  202. // now return an OAuth2 token without the id_token
  203. server.binding.OIDC.oauth2Config = &mockOAuth2Config{
  204. tokenSource: &mockTokenSource{},
  205. authCodeURL: webOIDCRedirectPath,
  206. token: &oauth2.Token{
  207. AccessToken: "123",
  208. Expiry: time.Now().Add(5 * time.Minute),
  209. },
  210. err: nil,
  211. }
  212. authReq := newOIDCPendingAuth(tokenAudienceWebClient)
  213. oidcMgr.addPendingAuth(authReq)
  214. rr = httptest.NewRecorder()
  215. r, err = http.NewRequest(http.MethodGet, webOIDCRedirectPath+"?state="+authReq.State, nil)
  216. assert.NoError(t, err)
  217. server.router.ServeHTTP(rr, r)
  218. assert.Equal(t, http.StatusFound, rr.Code)
  219. assert.Equal(t, webClientLoginPath, rr.Header().Get("Location"))
  220. require.Len(t, oidcMgr.pendingAuths, 0)
  221. // now fail to verify the id token
  222. token := &oauth2.Token{
  223. AccessToken: "123",
  224. Expiry: time.Now().Add(5 * time.Minute),
  225. }
  226. token = token.WithExtra(map[string]any{
  227. "id_token": "id_token_val",
  228. })
  229. server.binding.OIDC.oauth2Config = &mockOAuth2Config{
  230. tokenSource: &mockTokenSource{},
  231. authCodeURL: webOIDCRedirectPath,
  232. token: token,
  233. err: nil,
  234. }
  235. authReq = newOIDCPendingAuth(tokenAudienceWebClient)
  236. oidcMgr.addPendingAuth(authReq)
  237. rr = httptest.NewRecorder()
  238. r, err = http.NewRequest(http.MethodGet, webOIDCRedirectPath+"?state="+authReq.State, nil)
  239. assert.NoError(t, err)
  240. server.router.ServeHTTP(rr, r)
  241. assert.Equal(t, http.StatusFound, rr.Code)
  242. assert.Equal(t, webClientLoginPath, rr.Header().Get("Location"))
  243. require.Len(t, oidcMgr.pendingAuths, 0)
  244. // id token nonce does not match
  245. server.binding.OIDC.verifier = &mockOIDCVerifier{
  246. err: nil,
  247. token: &oidc.IDToken{},
  248. }
  249. authReq = newOIDCPendingAuth(tokenAudienceWebClient)
  250. oidcMgr.addPendingAuth(authReq)
  251. rr = httptest.NewRecorder()
  252. r, err = http.NewRequest(http.MethodGet, webOIDCRedirectPath+"?state="+authReq.State, nil)
  253. assert.NoError(t, err)
  254. server.router.ServeHTTP(rr, r)
  255. assert.Equal(t, http.StatusFound, rr.Code)
  256. assert.Equal(t, webClientLoginPath, rr.Header().Get("Location"))
  257. require.Len(t, oidcMgr.pendingAuths, 0)
  258. // null id token claims
  259. authReq = newOIDCPendingAuth(tokenAudienceWebClient)
  260. oidcMgr.addPendingAuth(authReq)
  261. server.binding.OIDC.verifier = &mockOIDCVerifier{
  262. err: nil,
  263. token: &oidc.IDToken{
  264. Nonce: authReq.Nonce,
  265. },
  266. }
  267. rr = httptest.NewRecorder()
  268. r, err = http.NewRequest(http.MethodGet, webOIDCRedirectPath+"?state="+authReq.State, nil)
  269. assert.NoError(t, err)
  270. server.router.ServeHTTP(rr, r)
  271. assert.Equal(t, http.StatusFound, rr.Code)
  272. assert.Equal(t, webClientLoginPath, rr.Header().Get("Location"))
  273. require.Len(t, oidcMgr.pendingAuths, 0)
  274. // invalid id token claims: no username
  275. authReq = newOIDCPendingAuth(tokenAudienceWebClient)
  276. oidcMgr.addPendingAuth(authReq)
  277. idToken := &oidc.IDToken{
  278. Nonce: authReq.Nonce,
  279. Expiry: time.Now().Add(5 * time.Minute),
  280. }
  281. setIDTokenClaims(idToken, []byte(`{"aud": "my_client_id"}`))
  282. server.binding.OIDC.verifier = &mockOIDCVerifier{
  283. err: nil,
  284. token: idToken,
  285. }
  286. rr = httptest.NewRecorder()
  287. r, err = http.NewRequest(http.MethodGet, webOIDCRedirectPath+"?state="+authReq.State, nil)
  288. assert.NoError(t, err)
  289. server.router.ServeHTTP(rr, r)
  290. assert.Equal(t, http.StatusFound, rr.Code)
  291. assert.Equal(t, webClientLoginPath, rr.Header().Get("Location"))
  292. require.Len(t, oidcMgr.pendingAuths, 0)
  293. // invalid id token clamims: username not a string
  294. authReq = newOIDCPendingAuth(tokenAudienceWebClient)
  295. oidcMgr.addPendingAuth(authReq)
  296. idToken = &oidc.IDToken{
  297. Nonce: authReq.Nonce,
  298. Expiry: time.Now().Add(5 * time.Minute),
  299. }
  300. setIDTokenClaims(idToken, []byte(`{"aud": "my_client_id","preferred_username": 1}`))
  301. server.binding.OIDC.verifier = &mockOIDCVerifier{
  302. err: nil,
  303. token: idToken,
  304. }
  305. rr = httptest.NewRecorder()
  306. r, err = http.NewRequest(http.MethodGet, webOIDCRedirectPath+"?state="+authReq.State, nil)
  307. assert.NoError(t, err)
  308. server.router.ServeHTTP(rr, r)
  309. assert.Equal(t, http.StatusFound, rr.Code)
  310. assert.Equal(t, webClientLoginPath, rr.Header().Get("Location"))
  311. require.Len(t, oidcMgr.pendingAuths, 0)
  312. // invalid audience
  313. authReq = newOIDCPendingAuth(tokenAudienceWebClient)
  314. oidcMgr.addPendingAuth(authReq)
  315. idToken = &oidc.IDToken{
  316. Nonce: authReq.Nonce,
  317. Expiry: time.Now().Add(5 * time.Minute),
  318. }
  319. setIDTokenClaims(idToken, []byte(`{"preferred_username":"test","sftpgo_role":"admin"}`))
  320. server.binding.OIDC.verifier = &mockOIDCVerifier{
  321. err: nil,
  322. token: idToken,
  323. }
  324. rr = httptest.NewRecorder()
  325. r, err = http.NewRequest(http.MethodGet, webOIDCRedirectPath+"?state="+authReq.State, nil)
  326. assert.NoError(t, err)
  327. server.router.ServeHTTP(rr, r)
  328. assert.Equal(t, http.StatusFound, rr.Code)
  329. assert.Equal(t, webClientLoginPath, rr.Header().Get("Location"))
  330. require.Len(t, oidcMgr.pendingAuths, 0)
  331. // invalid audience
  332. authReq = newOIDCPendingAuth(tokenAudienceWebAdmin)
  333. oidcMgr.addPendingAuth(authReq)
  334. idToken = &oidc.IDToken{
  335. Nonce: authReq.Nonce,
  336. Expiry: time.Now().Add(5 * time.Minute),
  337. }
  338. setIDTokenClaims(idToken, []byte(`{"preferred_username":"test"}`))
  339. server.binding.OIDC.verifier = &mockOIDCVerifier{
  340. err: nil,
  341. token: idToken,
  342. }
  343. rr = httptest.NewRecorder()
  344. r, err = http.NewRequest(http.MethodGet, webOIDCRedirectPath+"?state="+authReq.State, nil)
  345. assert.NoError(t, err)
  346. server.router.ServeHTTP(rr, r)
  347. assert.Equal(t, http.StatusFound, rr.Code)
  348. assert.Equal(t, webAdminLoginPath, rr.Header().Get("Location"))
  349. require.Len(t, oidcMgr.pendingAuths, 0)
  350. // mapped user not found
  351. authReq = newOIDCPendingAuth(tokenAudienceWebAdmin)
  352. oidcMgr.addPendingAuth(authReq)
  353. idToken = &oidc.IDToken{
  354. Nonce: authReq.Nonce,
  355. Expiry: time.Now().Add(5 * time.Minute),
  356. }
  357. setIDTokenClaims(idToken, []byte(`{"preferred_username":"test","sftpgo_role":"admin"}`))
  358. server.binding.OIDC.verifier = &mockOIDCVerifier{
  359. err: nil,
  360. token: idToken,
  361. }
  362. rr = httptest.NewRecorder()
  363. r, err = http.NewRequest(http.MethodGet, webOIDCRedirectPath+"?state="+authReq.State, nil)
  364. assert.NoError(t, err)
  365. server.router.ServeHTTP(rr, r)
  366. assert.Equal(t, http.StatusFound, rr.Code)
  367. assert.Equal(t, webAdminLoginPath, rr.Header().Get("Location"))
  368. require.Len(t, oidcMgr.pendingAuths, 0)
  369. // admin login ok
  370. authReq = newOIDCPendingAuth(tokenAudienceWebAdmin)
  371. oidcMgr.addPendingAuth(authReq)
  372. idToken = &oidc.IDToken{
  373. Nonce: authReq.Nonce,
  374. Expiry: time.Now().Add(5 * time.Minute),
  375. }
  376. setIDTokenClaims(idToken, []byte(`{"preferred_username":"admin","sftpgo_role":"admin","sid":"sid123"}`))
  377. server.binding.OIDC.verifier = &mockOIDCVerifier{
  378. err: nil,
  379. token: idToken,
  380. }
  381. rr = httptest.NewRecorder()
  382. r, err = http.NewRequest(http.MethodGet, webOIDCRedirectPath+"?state="+authReq.State, nil)
  383. assert.NoError(t, err)
  384. server.router.ServeHTTP(rr, r)
  385. assert.Equal(t, http.StatusFound, rr.Code)
  386. assert.Equal(t, webUsersPath, rr.Header().Get("Location"))
  387. require.Len(t, oidcMgr.pendingAuths, 0)
  388. require.Len(t, oidcMgr.tokens, 1)
  389. // admin profile is not available
  390. var tokenCookie string
  391. for k := range oidcMgr.tokens {
  392. tokenCookie = k
  393. }
  394. oidcToken, err := oidcMgr.getToken(tokenCookie)
  395. assert.NoError(t, err)
  396. assert.Equal(t, "sid123", oidcToken.SessionID)
  397. assert.True(t, oidcToken.isAdmin())
  398. assert.False(t, oidcToken.isExpired())
  399. rr = httptest.NewRecorder()
  400. r, err = http.NewRequest(http.MethodGet, webAdminProfilePath, nil)
  401. assert.NoError(t, err)
  402. r.Header.Set("Cookie", fmt.Sprintf("%v=%v", oidcCookieKey, tokenCookie))
  403. server.router.ServeHTTP(rr, r)
  404. assert.Equal(t, http.StatusForbidden, rr.Code)
  405. // the admin can access the allowed pages
  406. rr = httptest.NewRecorder()
  407. r, err = http.NewRequest(http.MethodGet, webUsersPath, nil)
  408. assert.NoError(t, err)
  409. r.Header.Set("Cookie", fmt.Sprintf("%v=%v", oidcCookieKey, tokenCookie))
  410. server.router.ServeHTTP(rr, r)
  411. assert.Equal(t, http.StatusOK, rr.Code)
  412. // try with an invalid cookie
  413. rr = httptest.NewRecorder()
  414. r, err = http.NewRequest(http.MethodGet, webUsersPath, nil)
  415. assert.NoError(t, err)
  416. r.Header.Set("Cookie", fmt.Sprintf("%v=%v", oidcCookieKey, xid.New().String()))
  417. server.router.ServeHTTP(rr, r)
  418. assert.Equal(t, http.StatusFound, rr.Code)
  419. assert.Equal(t, webAdminLoginPath, rr.Header().Get("Location"))
  420. // Web Client is not available with an admin token
  421. rr = httptest.NewRecorder()
  422. r, err = http.NewRequest(http.MethodGet, webClientFilesPath, nil)
  423. assert.NoError(t, err)
  424. r.Header.Set("Cookie", fmt.Sprintf("%v=%v", oidcCookieKey, tokenCookie))
  425. server.router.ServeHTTP(rr, r)
  426. assert.Equal(t, http.StatusFound, rr.Code)
  427. assert.Equal(t, webClientLoginPath, rr.Header().Get("Location"))
  428. // logout the admin user
  429. rr = httptest.NewRecorder()
  430. r, err = http.NewRequest(http.MethodGet, webLogoutPath, nil)
  431. assert.NoError(t, err)
  432. r.Header.Set("Cookie", fmt.Sprintf("%v=%v", oidcCookieKey, tokenCookie))
  433. server.router.ServeHTTP(rr, r)
  434. assert.Equal(t, http.StatusFound, rr.Code)
  435. assert.Equal(t, webAdminLoginPath, rr.Header().Get("Location"))
  436. require.Len(t, oidcMgr.pendingAuths, 0)
  437. require.Len(t, oidcMgr.tokens, 0)
  438. // now login and logout a user
  439. username := "test_oidc_user"
  440. user := dataprovider.User{
  441. BaseUser: sdk.BaseUser{
  442. Username: username,
  443. Password: "pwd",
  444. HomeDir: filepath.Join(os.TempDir(), username),
  445. Status: 1,
  446. Permissions: map[string][]string{
  447. "/": {dataprovider.PermAny},
  448. },
  449. },
  450. Filters: dataprovider.UserFilters{
  451. BaseUserFilters: sdk.BaseUserFilters{
  452. WebClient: []string{sdk.WebClientSharesDisabled},
  453. },
  454. },
  455. }
  456. err = dataprovider.AddUser(&user, "", "", "")
  457. assert.NoError(t, err)
  458. authReq = newOIDCPendingAuth(tokenAudienceWebClient)
  459. oidcMgr.addPendingAuth(authReq)
  460. idToken = &oidc.IDToken{
  461. Nonce: authReq.Nonce,
  462. Expiry: time.Now().Add(5 * time.Minute),
  463. }
  464. setIDTokenClaims(idToken, []byte(`{"preferred_username":"test_oidc_user"}`))
  465. server.binding.OIDC.verifier = &mockOIDCVerifier{
  466. err: nil,
  467. token: idToken,
  468. }
  469. rr = httptest.NewRecorder()
  470. r, err = http.NewRequest(http.MethodGet, webOIDCRedirectPath+"?state="+authReq.State, nil)
  471. assert.NoError(t, err)
  472. server.router.ServeHTTP(rr, r)
  473. assert.Equal(t, http.StatusFound, rr.Code)
  474. assert.Equal(t, webClientFilesPath, rr.Header().Get("Location"))
  475. require.Len(t, oidcMgr.pendingAuths, 0)
  476. require.Len(t, oidcMgr.tokens, 1)
  477. // user profile is not available
  478. for k := range oidcMgr.tokens {
  479. tokenCookie = k
  480. }
  481. oidcToken, err = oidcMgr.getToken(tokenCookie)
  482. assert.NoError(t, err)
  483. assert.Empty(t, oidcToken.SessionID)
  484. assert.False(t, oidcToken.isAdmin())
  485. assert.False(t, oidcToken.isExpired())
  486. if assert.Len(t, oidcToken.Permissions, 1) {
  487. assert.Equal(t, sdk.WebClientSharesDisabled, oidcToken.Permissions[0])
  488. }
  489. rr = httptest.NewRecorder()
  490. r, err = http.NewRequest(http.MethodGet, webClientProfilePath, nil)
  491. assert.NoError(t, err)
  492. r.RequestURI = webClientProfilePath
  493. r.Header.Set("Cookie", fmt.Sprintf("%v=%v", oidcCookieKey, tokenCookie))
  494. server.router.ServeHTTP(rr, r)
  495. assert.Equal(t, http.StatusOK, rr.Code)
  496. // the user can access the allowed pages
  497. rr = httptest.NewRecorder()
  498. r, err = http.NewRequest(http.MethodGet, webClientFilesPath, nil)
  499. assert.NoError(t, err)
  500. r.Header.Set("Cookie", fmt.Sprintf("%v=%v", oidcCookieKey, tokenCookie))
  501. server.router.ServeHTTP(rr, r)
  502. assert.Equal(t, http.StatusOK, rr.Code)
  503. // try with an invalid cookie
  504. rr = httptest.NewRecorder()
  505. r, err = http.NewRequest(http.MethodGet, webClientFilesPath, nil)
  506. assert.NoError(t, err)
  507. r.Header.Set("Cookie", fmt.Sprintf("%v=%v", oidcCookieKey, xid.New().String()))
  508. server.router.ServeHTTP(rr, r)
  509. assert.Equal(t, http.StatusFound, rr.Code)
  510. assert.Equal(t, webClientLoginPath, rr.Header().Get("Location"))
  511. // Web Admin is not available with a client cookie
  512. rr = httptest.NewRecorder()
  513. r, err = http.NewRequest(http.MethodGet, webUsersPath, nil)
  514. assert.NoError(t, err)
  515. r.Header.Set("Cookie", fmt.Sprintf("%v=%v", oidcCookieKey, tokenCookie))
  516. server.router.ServeHTTP(rr, r)
  517. assert.Equal(t, http.StatusFound, rr.Code)
  518. assert.Equal(t, webAdminLoginPath, rr.Header().Get("Location"))
  519. // logout the user
  520. rr = httptest.NewRecorder()
  521. r, err = http.NewRequest(http.MethodGet, webClientLogoutPath, nil)
  522. assert.NoError(t, err)
  523. r.Header.Set("Cookie", fmt.Sprintf("%v=%v", oidcCookieKey, tokenCookie))
  524. server.router.ServeHTTP(rr, r)
  525. assert.Equal(t, http.StatusFound, rr.Code)
  526. assert.Equal(t, webClientLoginPath, rr.Header().Get("Location"))
  527. require.Len(t, oidcMgr.pendingAuths, 0)
  528. require.Len(t, oidcMgr.tokens, 0)
  529. err = os.RemoveAll(user.GetHomeDir())
  530. assert.NoError(t, err)
  531. err = dataprovider.DeleteUser(username, "", "", "")
  532. assert.NoError(t, err)
  533. }
  534. func TestOIDCRefreshToken(t *testing.T) {
  535. oidcMgr, ok := oidcMgr.(*memoryOIDCManager)
  536. require.True(t, ok)
  537. r, err := http.NewRequest(http.MethodGet, webUsersPath, nil)
  538. assert.NoError(t, err)
  539. token := oidcToken{
  540. Cookie: xid.New().String(),
  541. AccessToken: xid.New().String(),
  542. TokenType: "Bearer",
  543. ExpiresAt: util.GetTimeAsMsSinceEpoch(time.Now().Add(-1 * time.Minute)),
  544. Nonce: xid.New().String(),
  545. Role: adminRoleFieldValue,
  546. Username: defaultAdminUsername,
  547. }
  548. config := mockOAuth2Config{
  549. tokenSource: &mockTokenSource{
  550. err: common.ErrGenericFailure,
  551. },
  552. }
  553. verifier := mockOIDCVerifier{
  554. err: common.ErrGenericFailure,
  555. }
  556. err = token.refresh(context.Background(), &config, &verifier, r)
  557. if assert.Error(t, err) {
  558. assert.Contains(t, err.Error(), "refresh token not set")
  559. }
  560. token.RefreshToken = xid.New().String()
  561. err = token.refresh(context.Background(), &config, &verifier, r)
  562. assert.ErrorIs(t, err, common.ErrGenericFailure)
  563. newToken := &oauth2.Token{
  564. AccessToken: xid.New().String(),
  565. RefreshToken: xid.New().String(),
  566. Expiry: time.Now().Add(5 * time.Minute),
  567. }
  568. config = mockOAuth2Config{
  569. tokenSource: &mockTokenSource{
  570. token: newToken,
  571. },
  572. }
  573. verifier = mockOIDCVerifier{
  574. token: &oidc.IDToken{},
  575. }
  576. err = token.refresh(context.Background(), &config, &verifier, r)
  577. if assert.Error(t, err) {
  578. assert.Contains(t, err.Error(), "the refreshed token has no id token")
  579. }
  580. newToken = newToken.WithExtra(map[string]any{
  581. "id_token": "id_token_val",
  582. })
  583. newToken.Expiry = time.Time{}
  584. config = mockOAuth2Config{
  585. tokenSource: &mockTokenSource{
  586. token: newToken,
  587. },
  588. }
  589. verifier = mockOIDCVerifier{
  590. err: common.ErrGenericFailure,
  591. }
  592. err = token.refresh(context.Background(), &config, &verifier, r)
  593. assert.ErrorIs(t, err, common.ErrGenericFailure)
  594. newToken = newToken.WithExtra(map[string]any{
  595. "id_token": "id_token_val",
  596. })
  597. newToken.Expiry = time.Now().Add(5 * time.Minute)
  598. config = mockOAuth2Config{
  599. tokenSource: &mockTokenSource{
  600. token: newToken,
  601. },
  602. }
  603. verifier = mockOIDCVerifier{
  604. token: &oidc.IDToken{
  605. Nonce: xid.New().String(), // nonce is different from the expected one
  606. },
  607. }
  608. err = token.refresh(context.Background(), &config, &verifier, r)
  609. if assert.Error(t, err) {
  610. assert.Contains(t, err.Error(), "the refreshed token nonce mismatch")
  611. }
  612. verifier = mockOIDCVerifier{
  613. token: &oidc.IDToken{
  614. Nonce: "", // empty token is fine on refresh but claims are not set
  615. },
  616. }
  617. err = token.refresh(context.Background(), &config, &verifier, r)
  618. if assert.Error(t, err) {
  619. assert.Contains(t, err.Error(), "oidc: claims not set")
  620. }
  621. idToken := &oidc.IDToken{
  622. Nonce: token.Nonce,
  623. }
  624. setIDTokenClaims(idToken, []byte(`{"sid":"id_token_sid"}`))
  625. verifier = mockOIDCVerifier{
  626. token: idToken,
  627. }
  628. err = token.refresh(context.Background(), &config, &verifier, r)
  629. assert.NoError(t, err)
  630. assert.Len(t, token.Permissions, 1)
  631. token.Role = nil
  632. // user does not exist
  633. err = token.refresh(context.Background(), &config, &verifier, r)
  634. assert.Error(t, err)
  635. require.Len(t, oidcMgr.tokens, 1)
  636. oidcMgr.removeToken(token.Cookie)
  637. require.Len(t, oidcMgr.tokens, 0)
  638. }
  639. func TestOIDCRefreshUser(t *testing.T) {
  640. token := oidcToken{
  641. Cookie: xid.New().String(),
  642. AccessToken: xid.New().String(),
  643. TokenType: "Bearer",
  644. ExpiresAt: util.GetTimeAsMsSinceEpoch(time.Now().Add(1 * time.Minute)),
  645. Nonce: xid.New().String(),
  646. Role: adminRoleFieldValue,
  647. Username: "missing username",
  648. }
  649. r, err := http.NewRequest(http.MethodGet, webUsersPath, nil)
  650. assert.NoError(t, err)
  651. err = token.refreshUser(r)
  652. assert.Error(t, err)
  653. admin := dataprovider.Admin{
  654. Username: "test_oidc_admin_refresh",
  655. Password: "p",
  656. Permissions: []string{dataprovider.PermAdminAny},
  657. Status: 0,
  658. Filters: dataprovider.AdminFilters{
  659. Preferences: dataprovider.AdminPreferences{
  660. HideUserPageSections: 1 + 2 + 4,
  661. },
  662. },
  663. }
  664. err = dataprovider.AddAdmin(&admin, "", "", "")
  665. assert.NoError(t, err)
  666. token.Username = admin.Username
  667. err = token.refreshUser(r)
  668. if assert.Error(t, err) {
  669. assert.Contains(t, err.Error(), "is disabled")
  670. }
  671. admin.Status = 1
  672. err = dataprovider.UpdateAdmin(&admin, "", "", "")
  673. assert.NoError(t, err)
  674. err = token.refreshUser(r)
  675. assert.NoError(t, err)
  676. assert.Equal(t, admin.Permissions, token.Permissions)
  677. assert.Equal(t, admin.Filters.Preferences.HideUserPageSections, token.HideUserPageSections)
  678. err = dataprovider.DeleteAdmin(admin.Username, "", "", "")
  679. assert.NoError(t, err)
  680. username := "test_oidc_user_refresh_token"
  681. user := dataprovider.User{
  682. BaseUser: sdk.BaseUser{
  683. Username: username,
  684. Password: "p",
  685. HomeDir: filepath.Join(os.TempDir(), username),
  686. Status: 0,
  687. Permissions: map[string][]string{
  688. "/": {dataprovider.PermAny},
  689. },
  690. },
  691. Filters: dataprovider.UserFilters{
  692. BaseUserFilters: sdk.BaseUserFilters{
  693. DeniedProtocols: []string{common.ProtocolHTTP},
  694. WebClient: []string{sdk.WebClientSharesDisabled, sdk.WebClientWriteDisabled},
  695. },
  696. },
  697. }
  698. err = dataprovider.AddUser(&user, "", "", "")
  699. assert.NoError(t, err)
  700. r, err = http.NewRequest(http.MethodGet, webClientFilesPath, nil)
  701. assert.NoError(t, err)
  702. token.Role = nil
  703. token.Username = username
  704. assert.False(t, token.isAdmin())
  705. err = token.refreshUser(r)
  706. if assert.Error(t, err) {
  707. assert.Contains(t, err.Error(), "is disabled")
  708. }
  709. user, err = dataprovider.UserExists(username, "")
  710. assert.NoError(t, err)
  711. user.Status = 1
  712. err = dataprovider.UpdateUser(&user, "", "", "")
  713. assert.NoError(t, err)
  714. err = token.refreshUser(r)
  715. if assert.Error(t, err) {
  716. assert.Contains(t, err.Error(), "protocol HTTP is not allowed")
  717. }
  718. user.Filters.DeniedProtocols = []string{common.ProtocolFTP}
  719. err = dataprovider.UpdateUser(&user, "", "", "")
  720. assert.NoError(t, err)
  721. err = token.refreshUser(r)
  722. assert.NoError(t, err)
  723. assert.Equal(t, user.Filters.WebClient, token.Permissions)
  724. err = dataprovider.DeleteUser(username, "", "", "")
  725. assert.NoError(t, err)
  726. }
  727. func TestValidateOIDCToken(t *testing.T) {
  728. oidcMgr, ok := oidcMgr.(*memoryOIDCManager)
  729. require.True(t, ok)
  730. server := getTestOIDCServer()
  731. err := server.binding.OIDC.initialize()
  732. assert.NoError(t, err)
  733. server.initializeRouter()
  734. rr := httptest.NewRecorder()
  735. r, err := http.NewRequest(http.MethodGet, webClientLogoutPath, nil)
  736. assert.NoError(t, err)
  737. _, err = server.validateOIDCToken(rr, r, false)
  738. assert.ErrorIs(t, err, errInvalidToken)
  739. // expired token and refresh error
  740. server.binding.OIDC.oauth2Config = &mockOAuth2Config{
  741. tokenSource: &mockTokenSource{
  742. err: common.ErrGenericFailure,
  743. },
  744. }
  745. token := oidcToken{
  746. Cookie: xid.New().String(),
  747. AccessToken: xid.New().String(),
  748. ExpiresAt: util.GetTimeAsMsSinceEpoch(time.Now().Add(-2 * time.Minute)),
  749. }
  750. oidcMgr.addToken(token)
  751. rr = httptest.NewRecorder()
  752. r, err = http.NewRequest(http.MethodGet, webClientLogoutPath, nil)
  753. assert.NoError(t, err)
  754. r.Header.Set("Cookie", fmt.Sprintf("%v=%v", oidcCookieKey, token.Cookie))
  755. _, err = server.validateOIDCToken(rr, r, false)
  756. assert.ErrorIs(t, err, errInvalidToken)
  757. oidcMgr.removeToken(token.Cookie)
  758. assert.Len(t, oidcMgr.tokens, 0)
  759. server.tokenAuth = jwtauth.New("PS256", util.GenerateRandomBytes(32), nil)
  760. token = oidcToken{
  761. Cookie: xid.New().String(),
  762. AccessToken: xid.New().String(),
  763. }
  764. oidcMgr.addToken(token)
  765. rr = httptest.NewRecorder()
  766. r, err = http.NewRequest(http.MethodGet, webClientLogoutPath, nil)
  767. assert.NoError(t, err)
  768. r.Header.Set("Cookie", fmt.Sprintf("%v=%v", oidcCookieKey, token.Cookie))
  769. server.router.ServeHTTP(rr, r)
  770. assert.Equal(t, http.StatusFound, rr.Code)
  771. assert.Equal(t, webClientLoginPath, rr.Header().Get("Location"))
  772. oidcMgr.removeToken(token.Cookie)
  773. assert.Len(t, oidcMgr.tokens, 0)
  774. token = oidcToken{
  775. Cookie: xid.New().String(),
  776. AccessToken: xid.New().String(),
  777. Role: "admin",
  778. }
  779. oidcMgr.addToken(token)
  780. rr = httptest.NewRecorder()
  781. r, err = http.NewRequest(http.MethodGet, webLogoutPath, nil)
  782. assert.NoError(t, err)
  783. r.Header.Set("Cookie", fmt.Sprintf("%v=%v", oidcCookieKey, token.Cookie))
  784. server.router.ServeHTTP(rr, r)
  785. assert.Equal(t, http.StatusFound, rr.Code)
  786. assert.Equal(t, webAdminLoginPath, rr.Header().Get("Location"))
  787. oidcMgr.removeToken(token.Cookie)
  788. assert.Len(t, oidcMgr.tokens, 0)
  789. }
  790. func TestSkipOIDCAuth(t *testing.T) {
  791. server := getTestOIDCServer()
  792. err := server.binding.OIDC.initialize()
  793. assert.NoError(t, err)
  794. server.initializeRouter()
  795. jwtTokenClaims := jwtTokenClaims{
  796. Username: "user",
  797. }
  798. _, tokenString, err := jwtTokenClaims.createToken(server.tokenAuth, tokenAudienceWebClient, "")
  799. assert.NoError(t, err)
  800. rr := httptest.NewRecorder()
  801. r, err := http.NewRequest(http.MethodGet, webClientLogoutPath, nil)
  802. assert.NoError(t, err)
  803. r.Header.Set("Cookie", fmt.Sprintf("%v=%v", jwtCookieKey, tokenString))
  804. server.router.ServeHTTP(rr, r)
  805. assert.Equal(t, http.StatusFound, rr.Code)
  806. assert.Equal(t, webClientLoginPath, rr.Header().Get("Location"))
  807. }
  808. func TestOIDCLogoutErrors(t *testing.T) {
  809. server := getTestOIDCServer()
  810. assert.Empty(t, server.binding.OIDC.providerLogoutURL)
  811. server.logoutFromOIDCOP("")
  812. server.binding.OIDC.providerLogoutURL = "http://foo\x7f.com/"
  813. server.doOIDCFromLogout("")
  814. server.binding.OIDC.providerLogoutURL = "http://127.0.0.1:11234"
  815. server.doOIDCFromLogout("")
  816. }
  817. func TestOIDCToken(t *testing.T) {
  818. admin := dataprovider.Admin{
  819. Username: "test_oidc_admin",
  820. Password: "p",
  821. Permissions: []string{dataprovider.PermAdminAny},
  822. Status: 0,
  823. }
  824. err := dataprovider.AddAdmin(&admin, "", "", "")
  825. assert.NoError(t, err)
  826. token := oidcToken{
  827. Username: admin.Username,
  828. }
  829. // role not initialized, user with the specified username does not exist
  830. req, err := http.NewRequest(http.MethodGet, webUsersPath, nil)
  831. assert.NoError(t, err)
  832. err = token.getUser(req)
  833. assert.ErrorIs(t, err, util.ErrNotFound)
  834. token.Role = "admin"
  835. req, err = http.NewRequest(http.MethodGet, webUsersPath, nil)
  836. assert.NoError(t, err)
  837. err = token.getUser(req)
  838. if assert.Error(t, err) {
  839. assert.Contains(t, err.Error(), "is disabled")
  840. }
  841. err = dataprovider.DeleteAdmin(admin.Username, "", "", "")
  842. assert.NoError(t, err)
  843. username := "test_oidc_user"
  844. token.Username = username
  845. token.Role = ""
  846. err = token.getUser(req)
  847. if assert.Error(t, err) {
  848. assert.ErrorIs(t, err, util.ErrNotFound)
  849. }
  850. user := dataprovider.User{
  851. BaseUser: sdk.BaseUser{
  852. Username: username,
  853. Password: "p",
  854. HomeDir: filepath.Join(os.TempDir(), username),
  855. Status: 0,
  856. Permissions: map[string][]string{
  857. "/": {dataprovider.PermAny},
  858. },
  859. },
  860. Filters: dataprovider.UserFilters{
  861. BaseUserFilters: sdk.BaseUserFilters{
  862. DeniedProtocols: []string{common.ProtocolHTTP},
  863. },
  864. },
  865. }
  866. err = dataprovider.AddUser(&user, "", "", "")
  867. assert.NoError(t, err)
  868. err = token.getUser(req)
  869. if assert.Error(t, err) {
  870. assert.Contains(t, err.Error(), "is disabled")
  871. }
  872. user, err = dataprovider.UserExists(username, "")
  873. assert.NoError(t, err)
  874. user.Status = 1
  875. user.Password = "np"
  876. err = dataprovider.UpdateUser(&user, "", "", "")
  877. assert.NoError(t, err)
  878. err = token.getUser(req)
  879. if assert.Error(t, err) {
  880. assert.Contains(t, err.Error(), "protocol HTTP is not allowed")
  881. }
  882. user.Filters.DeniedProtocols = nil
  883. user.FsConfig.Provider = sdk.SFTPFilesystemProvider
  884. user.FsConfig.SFTPConfig = vfs.SFTPFsConfig{
  885. BaseSFTPFsConfig: sdk.BaseSFTPFsConfig{
  886. Endpoint: "127.0.0.1:8022",
  887. Username: username,
  888. },
  889. Password: kms.NewPlainSecret("np"),
  890. }
  891. err = dataprovider.UpdateUser(&user, "", "", "")
  892. assert.NoError(t, err)
  893. err = token.getUser(req)
  894. if assert.Error(t, err) {
  895. assert.Contains(t, err.Error(), "SFTP loop")
  896. }
  897. common.Config.PostConnectHook = fmt.Sprintf("http://%v/404", oidcMockAddr)
  898. err = token.getUser(req)
  899. if assert.Error(t, err) {
  900. assert.Contains(t, err.Error(), "access denied")
  901. }
  902. common.Config.PostConnectHook = ""
  903. err = os.RemoveAll(user.GetHomeDir())
  904. assert.NoError(t, err)
  905. err = dataprovider.DeleteUser(username, "", "", "")
  906. assert.NoError(t, err)
  907. }
  908. func TestOIDCImplicitRoles(t *testing.T) {
  909. oidcMgr, ok := oidcMgr.(*memoryOIDCManager)
  910. require.True(t, ok)
  911. server := getTestOIDCServer()
  912. server.binding.OIDC.ImplicitRoles = true
  913. err := server.binding.OIDC.initialize()
  914. assert.NoError(t, err)
  915. server.initializeRouter()
  916. authReq := newOIDCPendingAuth(tokenAudienceWebAdmin)
  917. oidcMgr.addPendingAuth(authReq)
  918. token := &oauth2.Token{
  919. AccessToken: "1234",
  920. Expiry: time.Now().Add(5 * time.Minute),
  921. }
  922. token = token.WithExtra(map[string]any{
  923. "id_token": "id_token_val",
  924. })
  925. server.binding.OIDC.oauth2Config = &mockOAuth2Config{
  926. tokenSource: &mockTokenSource{},
  927. authCodeURL: webOIDCRedirectPath,
  928. token: token,
  929. }
  930. idToken := &oidc.IDToken{
  931. Nonce: authReq.Nonce,
  932. Expiry: time.Now().Add(5 * time.Minute),
  933. }
  934. setIDTokenClaims(idToken, []byte(`{"preferred_username":"admin","sid":"sid456"}`))
  935. server.binding.OIDC.verifier = &mockOIDCVerifier{
  936. err: nil,
  937. token: idToken,
  938. }
  939. rr := httptest.NewRecorder()
  940. r, err := http.NewRequest(http.MethodGet, webOIDCRedirectPath+"?state="+authReq.State, nil)
  941. assert.NoError(t, err)
  942. server.router.ServeHTTP(rr, r)
  943. assert.Equal(t, http.StatusFound, rr.Code)
  944. assert.Equal(t, webUsersPath, rr.Header().Get("Location"))
  945. require.Len(t, oidcMgr.pendingAuths, 0)
  946. require.Len(t, oidcMgr.tokens, 1)
  947. var tokenCookie string
  948. for k := range oidcMgr.tokens {
  949. tokenCookie = k
  950. }
  951. // Web Client is not available with an admin token
  952. rr = httptest.NewRecorder()
  953. r, err = http.NewRequest(http.MethodGet, webClientFilesPath, nil)
  954. assert.NoError(t, err)
  955. r.Header.Set("Cookie", fmt.Sprintf("%v=%v", oidcCookieKey, tokenCookie))
  956. server.router.ServeHTTP(rr, r)
  957. assert.Equal(t, http.StatusFound, rr.Code)
  958. assert.Equal(t, webClientLoginPath, rr.Header().Get("Location"))
  959. // logout the admin user
  960. rr = httptest.NewRecorder()
  961. r, err = http.NewRequest(http.MethodGet, webLogoutPath, nil)
  962. assert.NoError(t, err)
  963. r.Header.Set("Cookie", fmt.Sprintf("%v=%v", oidcCookieKey, tokenCookie))
  964. server.router.ServeHTTP(rr, r)
  965. assert.Equal(t, http.StatusFound, rr.Code)
  966. assert.Equal(t, webAdminLoginPath, rr.Header().Get("Location"))
  967. require.Len(t, oidcMgr.pendingAuths, 0)
  968. require.Len(t, oidcMgr.tokens, 0)
  969. // now login and logout a user
  970. username := "test_oidc_implicit_user"
  971. user := dataprovider.User{
  972. BaseUser: sdk.BaseUser{
  973. Username: username,
  974. Password: "pwd",
  975. HomeDir: filepath.Join(os.TempDir(), username),
  976. Status: 1,
  977. Permissions: map[string][]string{
  978. "/": {dataprovider.PermAny},
  979. },
  980. },
  981. Filters: dataprovider.UserFilters{
  982. BaseUserFilters: sdk.BaseUserFilters{
  983. WebClient: []string{sdk.WebClientSharesDisabled},
  984. },
  985. },
  986. }
  987. err = dataprovider.AddUser(&user, "", "", "")
  988. assert.NoError(t, err)
  989. authReq = newOIDCPendingAuth(tokenAudienceWebClient)
  990. oidcMgr.addPendingAuth(authReq)
  991. idToken = &oidc.IDToken{
  992. Nonce: authReq.Nonce,
  993. Expiry: time.Now().Add(5 * time.Minute),
  994. }
  995. setIDTokenClaims(idToken, []byte(`{"preferred_username":"test_oidc_implicit_user"}`))
  996. server.binding.OIDC.verifier = &mockOIDCVerifier{
  997. err: nil,
  998. token: idToken,
  999. }
  1000. rr = httptest.NewRecorder()
  1001. r, err = http.NewRequest(http.MethodGet, webOIDCRedirectPath+"?state="+authReq.State, nil)
  1002. assert.NoError(t, err)
  1003. server.router.ServeHTTP(rr, r)
  1004. assert.Equal(t, http.StatusFound, rr.Code)
  1005. assert.Equal(t, webClientFilesPath, rr.Header().Get("Location"))
  1006. require.Len(t, oidcMgr.pendingAuths, 0)
  1007. require.Len(t, oidcMgr.tokens, 1)
  1008. for k := range oidcMgr.tokens {
  1009. tokenCookie = k
  1010. }
  1011. rr = httptest.NewRecorder()
  1012. r, err = http.NewRequest(http.MethodGet, webClientLogoutPath, nil)
  1013. assert.NoError(t, err)
  1014. r.Header.Set("Cookie", fmt.Sprintf("%v=%v", oidcCookieKey, tokenCookie))
  1015. server.router.ServeHTTP(rr, r)
  1016. assert.Equal(t, http.StatusFound, rr.Code)
  1017. assert.Equal(t, webClientLoginPath, rr.Header().Get("Location"))
  1018. require.Len(t, oidcMgr.pendingAuths, 0)
  1019. require.Len(t, oidcMgr.tokens, 0)
  1020. err = os.RemoveAll(user.GetHomeDir())
  1021. assert.NoError(t, err)
  1022. err = dataprovider.DeleteUser(username, "", "", "")
  1023. assert.NoError(t, err)
  1024. }
  1025. func TestMemoryOIDCManager(t *testing.T) {
  1026. oidcMgr, ok := oidcMgr.(*memoryOIDCManager)
  1027. require.True(t, ok)
  1028. require.Len(t, oidcMgr.pendingAuths, 0)
  1029. authReq := newOIDCPendingAuth(tokenAudienceWebAdmin)
  1030. oidcMgr.addPendingAuth(authReq)
  1031. require.Len(t, oidcMgr.pendingAuths, 1)
  1032. _, err := oidcMgr.getPendingAuth(authReq.State)
  1033. assert.NoError(t, err)
  1034. oidcMgr.removePendingAuth(authReq.State)
  1035. require.Len(t, oidcMgr.pendingAuths, 0)
  1036. authReq.IssuedAt = util.GetTimeAsMsSinceEpoch(time.Now().Add(-61 * time.Second))
  1037. oidcMgr.addPendingAuth(authReq)
  1038. require.Len(t, oidcMgr.pendingAuths, 1)
  1039. _, err = oidcMgr.getPendingAuth(authReq.State)
  1040. if assert.Error(t, err) {
  1041. assert.Contains(t, err.Error(), "too old")
  1042. }
  1043. oidcMgr.cleanup()
  1044. require.Len(t, oidcMgr.pendingAuths, 0)
  1045. token := oidcToken{
  1046. AccessToken: xid.New().String(),
  1047. Nonce: xid.New().String(),
  1048. SessionID: xid.New().String(),
  1049. Cookie: xid.New().String(),
  1050. Username: xid.New().String(),
  1051. Role: "admin",
  1052. Permissions: []string{dataprovider.PermAdminAny},
  1053. }
  1054. require.Len(t, oidcMgr.tokens, 0)
  1055. oidcMgr.addToken(token)
  1056. require.Len(t, oidcMgr.tokens, 1)
  1057. _, err = oidcMgr.getToken(xid.New().String())
  1058. assert.Error(t, err)
  1059. storedToken, err := oidcMgr.getToken(token.Cookie)
  1060. assert.NoError(t, err)
  1061. token.UsedAt = 0 // ensure we don't modify the stored token
  1062. assert.Greater(t, storedToken.UsedAt, int64(0))
  1063. token.UsedAt = storedToken.UsedAt
  1064. assert.Equal(t, token, storedToken)
  1065. // the usage will not be updated, it is recent
  1066. oidcMgr.updateTokenUsage(storedToken)
  1067. storedToken, err = oidcMgr.getToken(token.Cookie)
  1068. assert.NoError(t, err)
  1069. assert.Equal(t, token, storedToken)
  1070. usedAt := util.GetTimeAsMsSinceEpoch(time.Now().Add(-5 * time.Minute))
  1071. storedToken.UsedAt = usedAt
  1072. oidcMgr.tokens[token.Cookie] = storedToken
  1073. storedToken, err = oidcMgr.getToken(token.Cookie)
  1074. assert.NoError(t, err)
  1075. assert.Equal(t, usedAt, storedToken.UsedAt)
  1076. token.UsedAt = storedToken.UsedAt
  1077. assert.Equal(t, token, storedToken)
  1078. oidcMgr.updateTokenUsage(storedToken)
  1079. storedToken, err = oidcMgr.getToken(token.Cookie)
  1080. assert.NoError(t, err)
  1081. assert.Greater(t, storedToken.UsedAt, usedAt)
  1082. token.UsedAt = storedToken.UsedAt
  1083. assert.Equal(t, token, storedToken)
  1084. storedToken.UsedAt = util.GetTimeAsMsSinceEpoch(time.Now()) - tokenDeleteInterval - 1
  1085. oidcMgr.tokens[token.Cookie] = storedToken
  1086. storedToken, err = oidcMgr.getToken(token.Cookie)
  1087. if assert.Error(t, err) {
  1088. assert.Contains(t, err.Error(), "token is too old")
  1089. }
  1090. oidcMgr.removeToken(xid.New().String())
  1091. require.Len(t, oidcMgr.tokens, 1)
  1092. oidcMgr.removeToken(token.Cookie)
  1093. require.Len(t, oidcMgr.tokens, 0)
  1094. oidcMgr.addToken(token)
  1095. usedAt = util.GetTimeAsMsSinceEpoch(time.Now().Add(-6 * time.Hour))
  1096. token.UsedAt = usedAt
  1097. oidcMgr.tokens[token.Cookie] = token
  1098. newToken := oidcToken{
  1099. Cookie: xid.New().String(),
  1100. }
  1101. oidcMgr.addToken(newToken)
  1102. oidcMgr.cleanup()
  1103. require.Len(t, oidcMgr.tokens, 1)
  1104. _, err = oidcMgr.getToken(token.Cookie)
  1105. assert.Error(t, err)
  1106. _, err = oidcMgr.getToken(newToken.Cookie)
  1107. assert.NoError(t, err)
  1108. oidcMgr.removeToken(newToken.Cookie)
  1109. require.Len(t, oidcMgr.tokens, 0)
  1110. }
  1111. func TestOIDCEvMgrIntegration(t *testing.T) {
  1112. providerConf := dataprovider.GetProviderConfig()
  1113. err := dataprovider.Close()
  1114. assert.NoError(t, err)
  1115. newProviderConf := providerConf
  1116. newProviderConf.NamingRules = 5
  1117. err = dataprovider.Initialize(newProviderConf, configDir, true)
  1118. assert.NoError(t, err)
  1119. // add a special chars to check json replacer
  1120. username := `test_"oidc_eventmanager`
  1121. u := map[string]any{
  1122. "username": "{{Name}}",
  1123. "status": 1,
  1124. "home_dir": filepath.Join(os.TempDir(), "{{IDPFieldcustom1.sub}}"),
  1125. "permissions": map[string][]string{
  1126. "/": {dataprovider.PermAny},
  1127. },
  1128. "description": "{{IDPFieldcustom2}}",
  1129. }
  1130. userTmpl, err := json.Marshal(u)
  1131. require.NoError(t, err)
  1132. a := map[string]any{
  1133. "username": "{{Name}}",
  1134. "status": 1,
  1135. "permissions": []string{dataprovider.PermAdminAny},
  1136. }
  1137. adminTmpl, err := json.Marshal(a)
  1138. require.NoError(t, err)
  1139. action := &dataprovider.BaseEventAction{
  1140. Name: "a",
  1141. Type: dataprovider.ActionTypeIDPAccountCheck,
  1142. Options: dataprovider.BaseEventActionOptions{
  1143. IDPConfig: dataprovider.EventActionIDPAccountCheck{
  1144. Mode: 0,
  1145. TemplateUser: string(userTmpl),
  1146. TemplateAdmin: string(adminTmpl),
  1147. },
  1148. },
  1149. }
  1150. err = dataprovider.AddEventAction(action, "", "", "")
  1151. assert.NoError(t, err)
  1152. rule := &dataprovider.EventRule{
  1153. Name: "r",
  1154. Status: 1,
  1155. Trigger: dataprovider.EventTriggerIDPLogin,
  1156. Conditions: dataprovider.EventConditions{
  1157. IDPLoginEvent: 0,
  1158. },
  1159. Actions: []dataprovider.EventAction{
  1160. {
  1161. BaseEventAction: dataprovider.BaseEventAction{
  1162. Name: action.Name,
  1163. },
  1164. Options: dataprovider.EventActionOptions{
  1165. ExecuteSync: true,
  1166. },
  1167. },
  1168. },
  1169. }
  1170. err = dataprovider.AddEventRule(rule, "", "", "")
  1171. assert.NoError(t, err)
  1172. oidcMgr, ok := oidcMgr.(*memoryOIDCManager)
  1173. require.True(t, ok)
  1174. server := getTestOIDCServer()
  1175. server.binding.OIDC.ImplicitRoles = true
  1176. server.binding.OIDC.CustomFields = []string{"custom1.sub", "custom2"}
  1177. err = server.binding.OIDC.initialize()
  1178. assert.NoError(t, err)
  1179. server.initializeRouter()
  1180. // login a user with OIDC
  1181. _, err = dataprovider.UserExists(username, "")
  1182. assert.ErrorIs(t, err, util.ErrNotFound)
  1183. authReq := newOIDCPendingAuth(tokenAudienceWebClient)
  1184. oidcMgr.addPendingAuth(authReq)
  1185. token := &oauth2.Token{
  1186. AccessToken: "1234",
  1187. Expiry: time.Now().Add(5 * time.Minute),
  1188. }
  1189. token = token.WithExtra(map[string]any{
  1190. "id_token": "id_token_val",
  1191. })
  1192. server.binding.OIDC.oauth2Config = &mockOAuth2Config{
  1193. tokenSource: &mockTokenSource{},
  1194. authCodeURL: webOIDCRedirectPath,
  1195. token: token,
  1196. }
  1197. idToken := &oidc.IDToken{
  1198. Nonce: authReq.Nonce,
  1199. Expiry: time.Now().Add(5 * time.Minute),
  1200. }
  1201. setIDTokenClaims(idToken, []byte(`{"preferred_username":"`+util.JSONEscape(username)+`","custom1":{"sub":"val1"},"custom2":"desc"}`)) //nolint:goconst
  1202. server.binding.OIDC.verifier = &mockOIDCVerifier{
  1203. err: nil,
  1204. token: idToken,
  1205. }
  1206. rr := httptest.NewRecorder()
  1207. r, err := http.NewRequest(http.MethodGet, webOIDCRedirectPath+"?state="+authReq.State, nil)
  1208. assert.NoError(t, err)
  1209. server.router.ServeHTTP(rr, r)
  1210. assert.Equal(t, http.StatusFound, rr.Code)
  1211. assert.Equal(t, webClientFilesPath, rr.Header().Get("Location"))
  1212. user, err := dataprovider.UserExists(username, "")
  1213. assert.NoError(t, err)
  1214. assert.Equal(t, filepath.Join(os.TempDir(), "val1"), user.GetHomeDir())
  1215. assert.Equal(t, "desc", user.Description)
  1216. err = dataprovider.DeleteUser(username, "", "", "")
  1217. assert.NoError(t, err)
  1218. err = os.RemoveAll(user.GetHomeDir())
  1219. assert.NoError(t, err)
  1220. // login an admin with OIDC
  1221. _, err = dataprovider.AdminExists(username)
  1222. assert.ErrorIs(t, err, util.ErrNotFound)
  1223. authReq = newOIDCPendingAuth(tokenAudienceWebAdmin)
  1224. oidcMgr.addPendingAuth(authReq)
  1225. idToken = &oidc.IDToken{
  1226. Nonce: authReq.Nonce,
  1227. Expiry: time.Now().Add(5 * time.Minute),
  1228. }
  1229. setIDTokenClaims(idToken, []byte(`{"preferred_username":"`+util.JSONEscape(username)+`"}`))
  1230. server.binding.OIDC.verifier = &mockOIDCVerifier{
  1231. err: nil,
  1232. token: idToken,
  1233. }
  1234. rr = httptest.NewRecorder()
  1235. r, err = http.NewRequest(http.MethodGet, webOIDCRedirectPath+"?state="+authReq.State, nil)
  1236. assert.NoError(t, err)
  1237. server.router.ServeHTTP(rr, r)
  1238. assert.Equal(t, http.StatusFound, rr.Code)
  1239. assert.Equal(t, webUsersPath, rr.Header().Get("Location"))
  1240. _, err = dataprovider.AdminExists(username)
  1241. assert.NoError(t, err)
  1242. err = dataprovider.DeleteAdmin(username, "", "", "")
  1243. assert.NoError(t, err)
  1244. // set invalid templates and try again
  1245. action.Options.IDPConfig.TemplateUser = `{}`
  1246. action.Options.IDPConfig.TemplateAdmin = `{}`
  1247. err = dataprovider.UpdateEventAction(action, "", "", "")
  1248. assert.NoError(t, err)
  1249. for _, audience := range []string{tokenAudienceWebAdmin, tokenAudienceWebClient} {
  1250. authReq = newOIDCPendingAuth(audience)
  1251. oidcMgr.addPendingAuth(authReq)
  1252. idToken = &oidc.IDToken{
  1253. Nonce: authReq.Nonce,
  1254. Expiry: time.Now().Add(5 * time.Minute),
  1255. }
  1256. setIDTokenClaims(idToken, []byte(`{"preferred_username":"`+util.JSONEscape(username)+`"}`))
  1257. server.binding.OIDC.verifier = &mockOIDCVerifier{
  1258. err: nil,
  1259. token: idToken,
  1260. }
  1261. rr = httptest.NewRecorder()
  1262. r, err = http.NewRequest(http.MethodGet, webOIDCRedirectPath+"?state="+authReq.State, nil)
  1263. assert.NoError(t, err)
  1264. server.router.ServeHTTP(rr, r)
  1265. assert.Equal(t, http.StatusFound, rr.Code)
  1266. }
  1267. for k := range oidcMgr.tokens {
  1268. oidcMgr.removeToken(k)
  1269. }
  1270. err = dataprovider.DeleteEventRule(rule.Name, "", "", "")
  1271. assert.NoError(t, err)
  1272. err = dataprovider.DeleteEventAction(action.Name, "", "", "")
  1273. assert.NoError(t, err)
  1274. err = dataprovider.Close()
  1275. assert.NoError(t, err)
  1276. err = dataprovider.Initialize(providerConf, configDir, true)
  1277. assert.NoError(t, err)
  1278. }
  1279. func TestOIDCPreLoginHook(t *testing.T) {
  1280. if runtime.GOOS == osWindows {
  1281. t.Skip("this test is not available on Windows")
  1282. }
  1283. oidcMgr, ok := oidcMgr.(*memoryOIDCManager)
  1284. require.True(t, ok)
  1285. username := "test_oidc_user_prelogin"
  1286. u := dataprovider.User{
  1287. BaseUser: sdk.BaseUser{
  1288. Username: username,
  1289. HomeDir: filepath.Join(os.TempDir(), username),
  1290. Status: 1,
  1291. Permissions: map[string][]string{
  1292. "/": {dataprovider.PermAny},
  1293. },
  1294. },
  1295. }
  1296. preLoginPath := filepath.Join(os.TempDir(), "prelogin.sh")
  1297. providerConf := dataprovider.GetProviderConfig()
  1298. err := dataprovider.Close()
  1299. assert.NoError(t, err)
  1300. err = os.WriteFile(preLoginPath, getPreLoginScriptContent(u, false), os.ModePerm)
  1301. assert.NoError(t, err)
  1302. newProviderConf := providerConf
  1303. newProviderConf.PreLoginHook = preLoginPath
  1304. err = dataprovider.Initialize(newProviderConf, configDir, true)
  1305. assert.NoError(t, err)
  1306. server := getTestOIDCServer()
  1307. server.binding.OIDC.CustomFields = []string{"field1", "field2"}
  1308. err = server.binding.OIDC.initialize()
  1309. assert.NoError(t, err)
  1310. server.initializeRouter()
  1311. _, err = dataprovider.UserExists(username, "")
  1312. assert.ErrorIs(t, err, util.ErrNotFound)
  1313. // now login with OIDC
  1314. authReq := newOIDCPendingAuth(tokenAudienceWebClient)
  1315. oidcMgr.addPendingAuth(authReq)
  1316. token := &oauth2.Token{
  1317. AccessToken: "1234",
  1318. Expiry: time.Now().Add(5 * time.Minute),
  1319. }
  1320. token = token.WithExtra(map[string]any{
  1321. "id_token": "id_token_val",
  1322. })
  1323. server.binding.OIDC.oauth2Config = &mockOAuth2Config{
  1324. tokenSource: &mockTokenSource{},
  1325. authCodeURL: webOIDCRedirectPath,
  1326. token: token,
  1327. }
  1328. idToken := &oidc.IDToken{
  1329. Nonce: authReq.Nonce,
  1330. Expiry: time.Now().Add(5 * time.Minute),
  1331. }
  1332. setIDTokenClaims(idToken, []byte(`{"preferred_username":"`+username+`"}`))
  1333. server.binding.OIDC.verifier = &mockOIDCVerifier{
  1334. err: nil,
  1335. token: idToken,
  1336. }
  1337. rr := httptest.NewRecorder()
  1338. r, err := http.NewRequest(http.MethodGet, webOIDCRedirectPath+"?state="+authReq.State, nil)
  1339. assert.NoError(t, err)
  1340. server.router.ServeHTTP(rr, r)
  1341. assert.Equal(t, http.StatusFound, rr.Code)
  1342. assert.Equal(t, webClientFilesPath, rr.Header().Get("Location"))
  1343. _, err = dataprovider.UserExists(username, "")
  1344. assert.NoError(t, err)
  1345. err = dataprovider.DeleteUser(username, "", "", "")
  1346. assert.NoError(t, err)
  1347. err = os.RemoveAll(u.HomeDir)
  1348. assert.NoError(t, err)
  1349. err = os.WriteFile(preLoginPath, getPreLoginScriptContent(u, true), os.ModePerm)
  1350. assert.NoError(t, err)
  1351. authReq = newOIDCPendingAuth(tokenAudienceWebClient)
  1352. oidcMgr.addPendingAuth(authReq)
  1353. idToken = &oidc.IDToken{
  1354. Nonce: authReq.Nonce,
  1355. Expiry: time.Now().Add(5 * time.Minute),
  1356. }
  1357. setIDTokenClaims(idToken, []byte(`{"preferred_username":"`+username+`","field1":"value1","field2":"value2","field3":"value3"}`))
  1358. server.binding.OIDC.verifier = &mockOIDCVerifier{
  1359. err: nil,
  1360. token: idToken,
  1361. }
  1362. rr = httptest.NewRecorder()
  1363. r, err = http.NewRequest(http.MethodGet, webOIDCRedirectPath+"?state="+authReq.State, nil)
  1364. assert.NoError(t, err)
  1365. server.router.ServeHTTP(rr, r)
  1366. assert.Equal(t, http.StatusFound, rr.Code)
  1367. assert.Equal(t, webClientLoginPath, rr.Header().Get("Location"))
  1368. _, err = dataprovider.UserExists(username, "")
  1369. assert.ErrorIs(t, err, util.ErrNotFound)
  1370. if assert.Len(t, oidcMgr.tokens, 1) {
  1371. for k := range oidcMgr.tokens {
  1372. oidcMgr.removeToken(k)
  1373. }
  1374. }
  1375. require.Len(t, oidcMgr.pendingAuths, 0)
  1376. require.Len(t, oidcMgr.tokens, 0)
  1377. err = dataprovider.Close()
  1378. assert.NoError(t, err)
  1379. err = dataprovider.Initialize(providerConf, configDir, true)
  1380. assert.NoError(t, err)
  1381. err = os.Remove(preLoginPath)
  1382. assert.NoError(t, err)
  1383. }
  1384. func TestOIDCIsAdmin(t *testing.T) {
  1385. type test struct {
  1386. input any
  1387. want bool
  1388. }
  1389. emptySlice := make([]any, 0)
  1390. tests := []test{
  1391. {input: "admin", want: true},
  1392. {input: append(emptySlice, "admin"), want: true},
  1393. {input: append(emptySlice, "user", "admin"), want: true},
  1394. {input: "user", want: false},
  1395. {input: emptySlice, want: false},
  1396. {input: append(emptySlice, 1), want: false},
  1397. {input: 1, want: false},
  1398. {input: nil, want: false},
  1399. {input: map[string]string{"admin": "admin"}, want: false},
  1400. }
  1401. for _, tc := range tests {
  1402. token := oidcToken{
  1403. Role: tc.input,
  1404. }
  1405. assert.Equal(t, tc.want, token.isAdmin(), "%v should return %t", tc.input, tc.want)
  1406. }
  1407. }
  1408. func TestParseAdminRole(t *testing.T) {
  1409. claims := make(map[string]any)
  1410. rawClaims := []byte(`{
  1411. "sub": "35666371",
  1412. "email": "[email protected]",
  1413. "preferred_username": "Sally",
  1414. "name": "Sally Tyler",
  1415. "updated_at": "2018-04-13T22:08:45Z",
  1416. "given_name": "Sally",
  1417. "family_name": "Tyler",
  1418. "params": {
  1419. "sftpgo_role": "admin",
  1420. "subparams": {
  1421. "sftpgo_role": "admin",
  1422. "inner": {
  1423. "sftpgo_role": ["user","admin"]
  1424. }
  1425. }
  1426. },
  1427. "at_hash": "lPLhxI2wjEndc-WfyroDZA",
  1428. "rt_hash": "mCmxPtA04N-55AxlEUbq-A",
  1429. "aud": "78d1d040-20c9-0136-5146-067351775fae92920",
  1430. "exp": 1523664997,
  1431. "iat": 1523657797
  1432. }`)
  1433. err := json.Unmarshal(rawClaims, &claims)
  1434. assert.NoError(t, err)
  1435. type test struct {
  1436. input string
  1437. want bool
  1438. val any
  1439. }
  1440. tests := []test{
  1441. {input: "", want: false},
  1442. {input: "sftpgo_role", want: false},
  1443. {input: "params.sftpgo_role", want: true, val: "admin"},
  1444. {input: "params.subparams.sftpgo_role", want: true, val: "admin"},
  1445. {input: "params.subparams.inner.sftpgo_role", want: true, val: []any{"user", "admin"}},
  1446. {input: "email", want: false},
  1447. {input: "missing", want: false},
  1448. {input: "params.email", want: false},
  1449. {input: "missing.sftpgo_role", want: false},
  1450. {input: "params", want: false},
  1451. {input: "params.subparams.inner.sftpgo_role.missing", want: false},
  1452. }
  1453. for _, tc := range tests {
  1454. token := oidcToken{}
  1455. token.getRoleFromField(claims, tc.input)
  1456. assert.Equal(t, tc.want, token.isAdmin(), "%q should return %t", tc.input, tc.want)
  1457. if tc.want {
  1458. assert.Equal(t, tc.val, token.Role)
  1459. }
  1460. }
  1461. }
  1462. func TestOIDCWithLoginFormsDisabled(t *testing.T) {
  1463. oidcMgr, ok := oidcMgr.(*memoryOIDCManager)
  1464. require.True(t, ok)
  1465. server := getTestOIDCServer()
  1466. server.binding.OIDC.ImplicitRoles = true
  1467. server.binding.EnabledLoginMethods = 3
  1468. server.binding.EnableWebAdmin = true
  1469. server.binding.EnableWebClient = true
  1470. err := server.binding.OIDC.initialize()
  1471. assert.NoError(t, err)
  1472. server.initializeRouter()
  1473. // login with an admin user
  1474. authReq := newOIDCPendingAuth(tokenAudienceWebAdmin)
  1475. oidcMgr.addPendingAuth(authReq)
  1476. token := &oauth2.Token{
  1477. AccessToken: "1234",
  1478. Expiry: time.Now().Add(5 * time.Minute),
  1479. }
  1480. token = token.WithExtra(map[string]any{
  1481. "id_token": "id_token_val",
  1482. })
  1483. server.binding.OIDC.oauth2Config = &mockOAuth2Config{
  1484. tokenSource: &mockTokenSource{},
  1485. authCodeURL: webOIDCRedirectPath,
  1486. token: token,
  1487. }
  1488. idToken := &oidc.IDToken{
  1489. Nonce: authReq.Nonce,
  1490. Expiry: time.Now().Add(5 * time.Minute),
  1491. }
  1492. setIDTokenClaims(idToken, []byte(`{"preferred_username":"admin","sid":"sid456"}`))
  1493. server.binding.OIDC.verifier = &mockOIDCVerifier{
  1494. err: nil,
  1495. token: idToken,
  1496. }
  1497. rr := httptest.NewRecorder()
  1498. r, err := http.NewRequest(http.MethodGet, webOIDCRedirectPath+"?state="+authReq.State, nil)
  1499. assert.NoError(t, err)
  1500. server.router.ServeHTTP(rr, r)
  1501. assert.Equal(t, http.StatusFound, rr.Code)
  1502. assert.Equal(t, webUsersPath, rr.Header().Get("Location"))
  1503. var tokenCookie string
  1504. for k := range oidcMgr.tokens {
  1505. tokenCookie = k
  1506. }
  1507. // we should be able to create admins without setting a password
  1508. adminUsername := "testAdmin"
  1509. form := make(url.Values)
  1510. form.Set(csrfFormToken, createCSRFToken(rr, r, server.csrfTokenAuth, tokenCookie, webBaseAdminPath))
  1511. form.Set("username", adminUsername)
  1512. form.Set("password", "")
  1513. form.Set("status", "1")
  1514. form.Set("permissions", "*")
  1515. rr = httptest.NewRecorder()
  1516. r, err = http.NewRequest(http.MethodPost, webAdminPath, bytes.NewBuffer([]byte(form.Encode())))
  1517. assert.NoError(t, err)
  1518. r.Header.Set("Cookie", fmt.Sprintf("%v=%v", oidcCookieKey, tokenCookie))
  1519. r.Header.Set("Content-Type", "application/x-www-form-urlencoded")
  1520. server.router.ServeHTTP(rr, r)
  1521. assert.Equal(t, http.StatusSeeOther, rr.Code)
  1522. _, err = dataprovider.AdminExists(adminUsername)
  1523. assert.NoError(t, err)
  1524. err = dataprovider.DeleteAdmin(adminUsername, "", "", "")
  1525. assert.NoError(t, err)
  1526. // login and password related routes are disabled
  1527. rr = httptest.NewRecorder()
  1528. r, err = http.NewRequest(http.MethodPost, webAdminLoginPath, nil)
  1529. assert.NoError(t, err)
  1530. server.router.ServeHTTP(rr, r)
  1531. assert.Equal(t, http.StatusMethodNotAllowed, rr.Code)
  1532. rr = httptest.NewRecorder()
  1533. r, err = http.NewRequest(http.MethodPost, webAdminTwoFactorPath, nil)
  1534. assert.NoError(t, err)
  1535. server.router.ServeHTTP(rr, r)
  1536. assert.Equal(t, http.StatusNotFound, rr.Code)
  1537. rr = httptest.NewRecorder()
  1538. r, err = http.NewRequest(http.MethodPost, webClientLoginPath, nil)
  1539. assert.NoError(t, err)
  1540. server.router.ServeHTTP(rr, r)
  1541. assert.Equal(t, http.StatusMethodNotAllowed, rr.Code)
  1542. rr = httptest.NewRecorder()
  1543. r, err = http.NewRequest(http.MethodPost, webClientForgotPwdPath, nil)
  1544. assert.NoError(t, err)
  1545. server.router.ServeHTTP(rr, r)
  1546. assert.Equal(t, http.StatusNotFound, rr.Code)
  1547. }
  1548. func TestDbOIDCManager(t *testing.T) {
  1549. if !isSharedProviderSupported() {
  1550. t.Skip("this test it is not available with this provider")
  1551. }
  1552. mgr := newOIDCManager(1)
  1553. pendingAuth := newOIDCPendingAuth(tokenAudienceWebAdmin)
  1554. mgr.addPendingAuth(pendingAuth)
  1555. authReq, err := mgr.getPendingAuth(pendingAuth.State)
  1556. assert.NoError(t, err)
  1557. assert.Equal(t, pendingAuth, authReq)
  1558. pendingAuth.IssuedAt = util.GetTimeAsMsSinceEpoch(time.Now().Add(-24 * time.Hour))
  1559. mgr.addPendingAuth(pendingAuth)
  1560. _, err = mgr.getPendingAuth(pendingAuth.State)
  1561. if assert.Error(t, err) {
  1562. assert.Contains(t, err.Error(), "auth request is too old")
  1563. }
  1564. mgr.removePendingAuth(pendingAuth.State)
  1565. _, err = mgr.getPendingAuth(pendingAuth.State)
  1566. if assert.Error(t, err) {
  1567. assert.Contains(t, err.Error(), "unable to get the auth request for the specified state")
  1568. }
  1569. mgr.addPendingAuth(pendingAuth)
  1570. _, err = mgr.getPendingAuth(pendingAuth.State)
  1571. if assert.Error(t, err) {
  1572. assert.Contains(t, err.Error(), "auth request is too old")
  1573. }
  1574. mgr.cleanup()
  1575. _, err = mgr.getPendingAuth(pendingAuth.State)
  1576. if assert.Error(t, err) {
  1577. assert.Contains(t, err.Error(), "unable to get the auth request for the specified state")
  1578. }
  1579. token := oidcToken{
  1580. Cookie: xid.New().String(),
  1581. AccessToken: xid.New().String(),
  1582. TokenType: "Bearer",
  1583. RefreshToken: xid.New().String(),
  1584. ExpiresAt: util.GetTimeAsMsSinceEpoch(time.Now().Add(-2 * time.Minute)),
  1585. SessionID: xid.New().String(),
  1586. IDToken: xid.New().String(),
  1587. Nonce: xid.New().String(),
  1588. Username: xid.New().String(),
  1589. Permissions: []string{dataprovider.PermAdminAny},
  1590. Role: "admin",
  1591. }
  1592. mgr.addToken(token)
  1593. tokenGet, err := mgr.getToken(token.Cookie)
  1594. assert.NoError(t, err)
  1595. assert.Greater(t, tokenGet.UsedAt, int64(0))
  1596. token.UsedAt = tokenGet.UsedAt
  1597. assert.Equal(t, token, tokenGet)
  1598. time.Sleep(100 * time.Millisecond)
  1599. mgr.updateTokenUsage(token)
  1600. // no change
  1601. tokenGet, err = mgr.getToken(token.Cookie)
  1602. assert.NoError(t, err)
  1603. assert.Equal(t, token.UsedAt, tokenGet.UsedAt)
  1604. tokenGet.UsedAt = util.GetTimeAsMsSinceEpoch(time.Now().Add(-24 * time.Hour))
  1605. tokenGet.RefreshToken = xid.New().String()
  1606. mgr.updateTokenUsage(tokenGet)
  1607. tokenGet, err = mgr.getToken(token.Cookie)
  1608. assert.NoError(t, err)
  1609. assert.NotEmpty(t, tokenGet.RefreshToken)
  1610. assert.NotEqual(t, token.RefreshToken, tokenGet.RefreshToken)
  1611. assert.Greater(t, tokenGet.UsedAt, token.UsedAt)
  1612. mgr.removeToken(token.Cookie)
  1613. tokenGet, err = mgr.getToken(token.Cookie)
  1614. if assert.Error(t, err) {
  1615. assert.Contains(t, err.Error(), "unable to get the token for the specified session")
  1616. }
  1617. // add an expired token
  1618. token.UsedAt = util.GetTimeAsMsSinceEpoch(time.Now().Add(-24 * time.Hour))
  1619. session := dataprovider.Session{
  1620. Key: token.Cookie,
  1621. Data: token,
  1622. Type: dataprovider.SessionTypeOIDCToken,
  1623. Timestamp: token.UsedAt + tokenDeleteInterval,
  1624. }
  1625. err = dataprovider.AddSharedSession(session)
  1626. assert.NoError(t, err)
  1627. _, err = mgr.getToken(token.Cookie)
  1628. if assert.Error(t, err) {
  1629. assert.Contains(t, err.Error(), "token is too old")
  1630. }
  1631. mgr.cleanup()
  1632. _, err = mgr.getToken(token.Cookie)
  1633. if assert.Error(t, err) {
  1634. assert.Contains(t, err.Error(), "unable to get the token for the specified session")
  1635. }
  1636. // adding a session without a key should fail
  1637. session.Key = ""
  1638. err = dataprovider.AddSharedSession(session)
  1639. if assert.Error(t, err) {
  1640. assert.Contains(t, err.Error(), "unable to save a session with an empty key")
  1641. }
  1642. session.Key = xid.New().String()
  1643. session.Type = 1000
  1644. err = dataprovider.AddSharedSession(session)
  1645. if assert.Error(t, err) {
  1646. assert.Contains(t, err.Error(), "invalid session type")
  1647. }
  1648. dbMgr, ok := mgr.(*dbOIDCManager)
  1649. if assert.True(t, ok) {
  1650. _, err = dbMgr.decodePendingAuthData(2)
  1651. assert.Error(t, err)
  1652. _, err = dbMgr.decodeTokenData(true)
  1653. assert.Error(t, err)
  1654. }
  1655. }
  1656. func getTestOIDCServer() *httpdServer {
  1657. return &httpdServer{
  1658. binding: Binding{
  1659. OIDC: OIDC{
  1660. ClientID: "sftpgo-client",
  1661. ClientSecret: "jRsmE0SWnuZjP7djBqNq0mrf8QN77j2c",
  1662. ConfigURL: fmt.Sprintf("http://%v/auth/realms/sftpgo", oidcMockAddr),
  1663. RedirectBaseURL: "http://127.0.0.1:8081/",
  1664. UsernameField: "preferred_username",
  1665. RoleField: "sftpgo_role",
  1666. ImplicitRoles: false,
  1667. Scopes: []string{oidc.ScopeOpenID, "profile", "email"},
  1668. CustomFields: nil,
  1669. Debug: true,
  1670. },
  1671. },
  1672. enableWebAdmin: true,
  1673. enableWebClient: true,
  1674. }
  1675. }
  1676. func getPreLoginScriptContent(user dataprovider.User, nonJSONResponse bool) []byte {
  1677. content := []byte("#!/bin/sh\n\n")
  1678. if nonJSONResponse {
  1679. content = append(content, []byte("echo 'text response'\n")...)
  1680. return content
  1681. }
  1682. if len(user.Username) > 0 {
  1683. u, _ := json.Marshal(user)
  1684. content = append(content, []byte(fmt.Sprintf("echo '%v'\n", string(u)))...)
  1685. }
  1686. return content
  1687. }