diff --git a/backend/internal/dto/oidc_dto.go b/backend/internal/dto/oidc_dto.go index 18063016..53249868 100644 --- a/backend/internal/dto/oidc_dto.go +++ b/backend/internal/dto/oidc_dto.go @@ -55,8 +55,8 @@ type OidcClientUpdateDto struct { LogoURL *string `json:"logoUrl"` DarkLogoURL *string `json:"darkLogoUrl"` IsGroupRestricted bool `json:"isGroupRestricted"` - AccessTokenDurationMinutes int64 `json:"accessTokenDurationMinutes" binding:"required,token_duration"` - RefreshTokenDurationMinutes int64 `json:"refreshTokenDurationMinutes" binding:"required,token_duration"` + AccessTokenDurationMinutes int64 `json:"accessTokenDurationMinutes" binding:"omitempty,token_duration"` + RefreshTokenDurationMinutes int64 `json:"refreshTokenDurationMinutes" binding:"omitempty,token_duration"` } type OidcClientCreateDto struct { diff --git a/backend/internal/dto/oidc_dto_test.go b/backend/internal/dto/oidc_dto_test.go new file mode 100644 index 00000000..4d257f5c --- /dev/null +++ b/backend/internal/dto/oidc_dto_test.go @@ -0,0 +1,72 @@ +package dto + +import ( + "encoding/json" + "testing" + + "github.com/gin-gonic/gin/binding" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestOidcClientUpdateDto_tokenLifetimes(t *testing.T) { + const baseFields = `"name":"Test Client","callbackURLs":["https://example.com/callback"]` + + for _, test := range []struct { + name string + body string + wantErr bool + wantAccess int64 + wantRefresh int64 + }{ + { + name: "lifetimes omitted", + body: `{` + baseFields + `}`, + }, + { + name: "lifetimes null", + body: `{` + baseFields + `,"accessTokenDurationMinutes":null,"refreshTokenDurationMinutes":null}`, + }, + { + name: "lifetimes zero", + body: `{` + baseFields + `,"accessTokenDurationMinutes":0,"refreshTokenDurationMinutes":0}`, + }, + { + name: "only the access token lifetime set", + body: `{` + baseFields + `,"accessTokenDurationMinutes":120}`, + wantAccess: 120, + }, + { + name: "both lifetimes set", + body: `{` + baseFields + `,"accessTokenDurationMinutes":120,"refreshTokenDurationMinutes":10080}`, + wantAccess: 120, + wantRefresh: 10080, + }, + { + name: "negative value is rejected", + body: `{` + baseFields + `,"accessTokenDurationMinutes":-1}`, + wantErr: true, + }, + { + name: "value above the maximum is rejected", + body: `{` + baseFields + `,"refreshTokenDurationMinutes":525601}`, + wantErr: true, + }, + } { + t.Run(test.name, func(t *testing.T) { + var input OidcClientUpdateDto + require.NoError(t, json.Unmarshal([]byte(test.body), &input)) + + err := binding.Validator.ValidateStruct(input) + if test.wantErr { + require.Error(t, err) + return + } + require.NoError(t, err) + + // A zero value is what the service turns into the default lifetime + assert.Equal(t, test.wantAccess, input.AccessTokenDurationMinutes) + assert.Equal(t, test.wantRefresh, input.RefreshTokenDurationMinutes) + }) + } +} diff --git a/backend/internal/dto/validations_test.go b/backend/internal/dto/validations_test.go index a3d7bfd0..59ba8fdb 100644 --- a/backend/internal/dto/validations_test.go +++ b/backend/internal/dto/validations_test.go @@ -10,7 +10,7 @@ import ( func TestTokenDurationValidation(t *testing.T) { type input struct { - Duration int64 `binding:"required,token_duration"` + Duration int64 `binding:"omitempty,token_duration"` } for _, test := range []struct { @@ -18,10 +18,11 @@ func TestTokenDurationValidation(t *testing.T) { value int64 wantErr bool }{ - {name: "omitted", wantErr: true}, - {name: "below minimum", value: 0, wantErr: true}, + {name: "omitted (default)"}, + {name: "negative", value: -1, wantErr: true}, {name: "minimum", value: 1}, {name: "custom duration", value: 90}, + {name: "maximum", value: 365 * 24 * 60}, {name: "above maximum", value: 365*24*60 + 1, wantErr: true}, } { t.Run(test.name, func(t *testing.T) { diff --git a/backend/internal/service/oidc_service.go b/backend/internal/service/oidc_service.go index a0bad26d..9498ee16 100644 --- a/backend/internal/service/oidc_service.go +++ b/backend/internal/service/oidc_service.go @@ -1,6 +1,7 @@ package service import ( + "cmp" "context" "errors" "fmt" @@ -148,9 +149,7 @@ func (s *OidcService) CreateClient(ctx context.Context, input dto.OidcClientCrea Base: model.Base{ ID: input.ID, }, - CreatedByID: new(userID), - AccessTokenDurationMinutes: model.DefaultAccessTokenDurationMinutes, - RefreshTokenDurationMinutes: model.DefaultRefreshTokenDurationMinutes, + CreatedByID: new(userID), } updateOIDCClientModelFromDto(&client, &input.OidcClientUpdateDto) @@ -257,8 +256,10 @@ func updateOIDCClientModelFromDto(client *model.OidcClient, input *dto.OidcClien client.SkipConsent = input.SkipConsent client.LaunchURL = input.LaunchURL client.IsGroupRestricted = input.IsGroupRestricted - client.AccessTokenDurationMinutes = input.AccessTokenDurationMinutes - client.RefreshTokenDurationMinutes = input.RefreshTokenDurationMinutes + + // Token lifetimes are optional, so a zero value falls back to the default + client.AccessTokenDurationMinutes = cmp.Or(input.AccessTokenDurationMinutes, model.DefaultAccessTokenDurationMinutes) + client.RefreshTokenDurationMinutes = cmp.Or(input.RefreshTokenDurationMinutes, model.DefaultRefreshTokenDurationMinutes) // Preserve fields that are sourced from the client metadata document if client.IsMetadataDocument() { diff --git a/backend/internal/service/oidc_service_test.go b/backend/internal/service/oidc_service_test.go index 786b454c..73fdd0e2 100644 --- a/backend/internal/service/oidc_service_test.go +++ b/backend/internal/service/oidc_service_test.go @@ -528,11 +528,9 @@ func TestOidcService_CreateClient_withDescription(t *testing.T) { description := "A test client description" input := dto.OidcClientCreateDto{ OidcClientUpdateDto: dto.OidcClientUpdateDto{ - Name: "Test Client", - Description: description, - CallbackURLs: []string{"https://example.com/callback"}, - AccessTokenDurationMinutes: model.DefaultAccessTokenDurationMinutes, - RefreshTokenDurationMinutes: model.DefaultRefreshTokenDurationMinutes, + Name: "Test Client", + Description: description, + CallbackURLs: []string{"https://example.com/callback"}, }, } @@ -554,10 +552,8 @@ func TestOidcService_CreateClient_withoutDescription(t *testing.T) { input := dto.OidcClientCreateDto{ OidcClientUpdateDto: dto.OidcClientUpdateDto{ - Name: "Test Client", - CallbackURLs: []string{"https://example.com/callback"}, - AccessTokenDurationMinutes: model.DefaultAccessTokenDurationMinutes, - RefreshTokenDurationMinutes: model.DefaultRefreshTokenDurationMinutes, + Name: "Test Client", + CallbackURLs: []string{"https://example.com/callback"}, }, } @@ -570,6 +566,97 @@ func TestOidcService_CreateClient_withoutDescription(t *testing.T) { assert.Empty(t, fetched.Description) } +func TestOidcService_CreateClient_tokenLifetimes(t *testing.T) { + for _, test := range []struct { + name string + access int64 + refresh int64 + wantAccess int64 + wantRefresh int64 + }{ + { + name: "omitted lifetimes fall back to the defaults", + wantAccess: model.DefaultAccessTokenDurationMinutes, + wantRefresh: model.DefaultRefreshTokenDurationMinutes, + }, + { + name: "only one lifetime provided", + access: 2 * 60, + wantAccess: 2 * 60, + wantRefresh: model.DefaultRefreshTokenDurationMinutes, + }, + { + name: "both lifetimes provided", + access: 2 * 60, + refresh: 7 * 24 * 60, + wantAccess: 2 * 60, + wantRefresh: 7 * 24 * 60, + }, + } { + t.Run(test.name, func(t *testing.T) { + db := testutils.NewDatabaseForTest(t) + + s, err := NewOidcService(db, nil, nil, nil, nil, nil, nil) + require.NoError(t, err) + + input := dto.OidcClientCreateDto{ + OidcClientUpdateDto: dto.OidcClientUpdateDto{ + Name: "Test Client", + CallbackURLs: []string{"https://example.com/callback"}, + AccessTokenDurationMinutes: test.access, + RefreshTokenDurationMinutes: test.refresh, + }, + } + + client, err := s.CreateClient(t.Context(), input, "user-id") + require.NoError(t, err) + + var fetched model.OidcClient + err = db.First(&fetched, "id = ?", client.ID).Error + require.NoError(t, err) + assert.Equal(t, test.wantAccess, fetched.AccessTokenDurationMinutes) + assert.Equal(t, test.wantRefresh, fetched.RefreshTokenDurationMinutes) + }) + } +} + +func TestOidcService_UpdateClient_tokenLifetimes(t *testing.T) { + db := testutils.NewDatabaseForTest(t) + + s, err := NewOidcService(db, nil, nil, nil, nil, nil, nil) + require.NoError(t, err) + + client := model.OidcClient{ + Name: "Test Client", + CallbackURLs: datatype.StringList{"https://example.com/callback"}, + AccessTokenDurationMinutes: 2 * 60, + RefreshTokenDurationMinutes: 7 * 24 * 60, + } + require.NoError(t, db.Create(&client).Error) + + // A request that predates the token lifetime fields must be accepted, and resets them to the defaults + input := dto.OidcClientUpdateDto{ + Name: "Test Client", + CallbackURLs: []string{"https://example.com/callback"}, + } + _, err = s.UpdateClient(t.Context(), client.ID, input) + require.NoError(t, err) + + var fetched model.OidcClient + require.NoError(t, db.First(&fetched, "id = ?", client.ID).Error) + assert.Equal(t, model.DefaultAccessTokenDurationMinutes, fetched.AccessTokenDurationMinutes) + assert.Equal(t, model.DefaultRefreshTokenDurationMinutes, fetched.RefreshTokenDurationMinutes) + + // Providing only one of them leaves the other at its default + input.AccessTokenDurationMinutes = 3 * 60 + _, err = s.UpdateClient(t.Context(), client.ID, input) + require.NoError(t, err) + + require.NoError(t, db.First(&fetched, "id = ?", client.ID).Error) + assert.Equal(t, int64(3*60), fetched.AccessTokenDurationMinutes) + assert.Equal(t, model.DefaultRefreshTokenDurationMinutes, fetched.RefreshTokenDurationMinutes) +} + func TestOidcService_CreateClientSecret_withCustomSecret(t *testing.T) { db := testutils.NewDatabaseForTest(t) @@ -610,11 +697,9 @@ func TestOidcService_UpdateClient_description(t *testing.T) { // Update with a description description := "Updated description" input := dto.OidcClientUpdateDto{ - Name: "Test Client", - Description: description, - CallbackURLs: []string{"https://example.com/callback"}, - AccessTokenDurationMinutes: model.DefaultAccessTokenDurationMinutes, - RefreshTokenDurationMinutes: model.DefaultRefreshTokenDurationMinutes, + Name: "Test Client", + Description: description, + CallbackURLs: []string{"https://example.com/callback"}, } _, err = s.UpdateClient(t.Context(), client.ID, input) @@ -728,9 +813,7 @@ func TestOidcService_UpdateClient_CIMDDoesNotOverwriteConcurrentMetadataRefresh( `).Error) input := dto.OidcClientUpdateDto{ - Description: "Locally managed description", - AccessTokenDurationMinutes: model.DefaultAccessTokenDurationMinutes, - RefreshTokenDurationMinutes: model.DefaultRefreshTokenDurationMinutes, + Description: "Locally managed description", } _, err = s.UpdateClient(t.Context(), client.ID, input) require.NoError(t, err) diff --git a/tests/specs/oidc-client-settings.spec.ts b/tests/specs/oidc-client-settings.spec.ts index 49caff98..af4fe6d0 100644 --- a/tests/specs/oidc-client-settings.spec.ts +++ b/tests/specs/oidc-client-settings.spec.ts @@ -2,11 +2,6 @@ import test, { expect, Page } from '@playwright/test'; import { oidcClients, userGroups } from '../data'; import { cleanupBackend } from '../utils/cleanup.util'; -const defaultTokenLifetimes = { - accessTokenDurationMinutes: 60, - refreshTokenDurationMinutes: 30 * 24 * 60 -}; - test.beforeEach(async () => await cleanupBackend()); test.describe('Create OIDC client', () => { @@ -246,7 +241,6 @@ test('Filter OIDC clients by PAR requirement', async ({ page, request }) => { // Enable PAR on the PAR test client await request.put(`/api/oidc/clients/${parClient.id}`, { data: { - ...defaultTokenLifetimes, name: parClient.name, callbackURLs: [parClient.callbackUrl], logoutCallbackURLs: [], diff --git a/tests/specs/oidc.spec.ts b/tests/specs/oidc.spec.ts index 08f64ea5..ecd32619 100644 --- a/tests/specs/oidc.spec.ts +++ b/tests/specs/oidc.spec.ts @@ -5,11 +5,6 @@ import { generateIdToken } from '../utils/jwt.util'; import * as oidcUtil from '../utils/oidc.util'; import passkeyUtil from '../utils/passkey.util'; -const defaultTokenLifetimes = { - accessTokenDurationMinutes: 60, - refreshTokenDurationMinutes: 30 * 24 * 60 -}; - test.beforeEach(async () => await cleanupBackend()); async function generateSeededOauthAccessToken( @@ -765,7 +760,6 @@ test('Device authorization flow forces reauthentication when client requires it' const client = oidcClients.nextcloud; await request.put(`/api/oidc/clients/${client.id}`, { data: { - ...defaultTokenLifetimes, name: client.name, callbackURLs: [client.callbackUrl], logoutCallbackURLs: [client.logoutCallbackUrl], @@ -891,7 +885,6 @@ test('Forces reauthentication when client requires it', async ({ page, request } await request.put(`/api/oidc/clients/${oidcClients.nextcloud.id}`, { data: { - ...defaultTokenLifetimes, name: oidcClients.nextcloud.name, callbackURLs: [oidcClients.nextcloud.callbackUrl], logoutCallbackURLs: [oidcClients.nextcloud.logoutCallbackUrl], @@ -1445,7 +1438,6 @@ test.describe('Pushed Authorization Requests (PAR)', () => { await page.request.put(`/api/oidc/clients/${client.id}`, { headers: { 'Content-Type': 'application/json' }, data: { - ...defaultTokenLifetimes, name: client.name, callbackURLs: [client.callbackUrl], logoutCallbackURLs: [], @@ -1489,7 +1481,6 @@ test.describe('Pushed Authorization Requests (PAR)', () => { await request.put(`/api/oidc/clients/${client.id}`, { headers: { 'Content-Type': 'application/json' }, data: { - ...defaultTokenLifetimes, name: client.name, callbackURLs: [client.callbackUrl], logoutCallbackURLs: [],