Files
pocket-id/backend/internal/service/oidc_service_test.go
T

1035 lines
35 KiB
Go

package service
import (
"io"
"net/http"
"strconv"
"strings"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/pocket-id/pocket-id/backend/internal/apperror"
"github.com/pocket-id/pocket-id/backend/internal/dto"
"github.com/pocket-id/pocket-id/backend/internal/model"
datatype "github.com/pocket-id/pocket-id/backend/internal/model/types"
"github.com/pocket-id/pocket-id/backend/internal/oidc"
"github.com/pocket-id/pocket-id/backend/internal/storage"
"github.com/pocket-id/pocket-id/backend/internal/utils"
testutils "github.com/pocket-id/pocket-id/backend/internal/utils/testing"
)
func TestListAuthorizedClientsRejectsMissingUser(t *testing.T) {
service := &OidcService{db: testutils.NewDatabaseForTest(t)}
_, _, err := service.ListAuthorizedClients(t.Context(), "missing-user", utils.ListRequestOptions{})
require.True(t, apperror.IsCode(err, apperror.CodeUserNotFound))
}
func TestOidcService_DeleteClientDeletesOAuth2Sessions(t *testing.T) {
db := testutils.NewDatabaseForTest(t)
require.NoError(t, db.Exec("PRAGMA foreign_keys = ON").Error)
client := model.OidcClient{Base: model.Base{ID: "deleted-client"}, Name: "Deleted Client"}
otherClient := model.OidcClient{Base: model.Base{ID: "other-client"}, Name: "Other Client"}
require.NoError(t, db.Create(&client).Error)
require.NoError(t, db.Create(&otherClient).Error)
for i, kind := range []string{"authorize_code", "access_token", "refresh_token", "par", "device_code"} {
session := oidc.OAuth2Session{
Base: model.Base{ID: "deleted-client-session-" + strconv.Itoa(i)},
Kind: kind,
Key: "deleted-client-key-" + strconv.Itoa(i),
RequestID: "deleted-client-request",
ClientID: client.ID,
Active: true,
RequestData: `{"client_id":"deleted-client","session":{"subject":"test-user","id_token_claims":{"jti":"test-jti"}}}`,
}
require.NoError(t, db.Create(&session).Error)
}
require.NoError(t, db.Create(&oidc.OAuth2Session{
Base: model.Base{ID: "other-client-session"},
Kind: "refresh_token",
Key: "other-client-key",
RequestID: "other-client-request",
ClientID: otherClient.ID,
Active: true,
RequestData: `{"client_id":"other-client","session":{"subject":"test-user"}}`,
}).Error)
service := &OidcService{db: db}
require.NoError(t, service.DeleteClient(t.Context(), client.ID))
var deletedClientSessionCount int64
require.NoError(t, db.Model(&oidc.OAuth2Session{}).Where("client_id = ?", client.ID).Count(&deletedClientSessionCount).Error)
assert.Zero(t, deletedClientSessionCount)
var otherClientSessionCount int64
require.NoError(t, db.Model(&oidc.OAuth2Session{}).Where("client_id = ?", otherClient.ID).Count(&otherClientSessionCount).Error)
assert.Equal(t, int64(1), otherClientSessionCount)
}
func TestOidcService_updateClientLogoType(t *testing.T) {
// Create a test database
db := testutils.NewDatabaseForTest(t)
// Create database storage
dbStorage, err := storage.NewDatabaseStorage(db)
require.NoError(t, err)
// Init the OidcService
s := &OidcService{
db: db,
fileStorage: dbStorage,
}
// Create a test client
client := model.OidcClient{
Name: "Test Client",
CallbackURLs: datatype.StringList{"https://example.com/callback"},
}
err = db.Create(&client).Error
require.NoError(t, err)
// Helper function to check if a file exists in storage
fileExists := func(t *testing.T, path string) bool {
t.Helper()
_, _, err := dbStorage.Open(t.Context(), path)
return err == nil
}
// Helper function to create a dummy file in storage
createDummyFile := func(t *testing.T, path string) {
t.Helper()
err := dbStorage.Save(t.Context(), path, strings.NewReader("dummy content"))
require.NoError(t, err)
}
t.Run("Updates light logo type for client without previous logo", func(t *testing.T) {
// Update the logo type
err := s.updateClientLogoType(t.Context(), client.ID, "png", true)
require.NoError(t, err)
// Verify the client was updated
var updatedClient model.OidcClient
err = db.First(&updatedClient, "id = ?", client.ID).Error
require.NoError(t, err)
require.NotNil(t, updatedClient.ImageType)
assert.Equal(t, "png", *updatedClient.ImageType)
})
t.Run("Updates dark logo type for client without previous dark logo", func(t *testing.T) {
// Update the dark logo type
err := s.updateClientLogoType(t.Context(), client.ID, "jpg", false)
require.NoError(t, err)
// Verify the client was updated
var updatedClient model.OidcClient
err = db.First(&updatedClient, "id = ?", client.ID).Error
require.NoError(t, err)
require.NotNil(t, updatedClient.DarkImageType)
assert.Equal(t, "jpg", *updatedClient.DarkImageType)
})
t.Run("Updates light logo type and deletes old file when type changes", func(t *testing.T) {
// Create the old PNG file in storage
oldPath := "oidc-client-images/" + client.ID + ".png"
createDummyFile(t, oldPath)
require.True(t, fileExists(t, oldPath), "Old file should exist before update")
// Client currently has a PNG logo, update to WEBP
err := s.updateClientLogoType(t.Context(), client.ID, "webp", true)
require.NoError(t, err)
// Verify the client was updated
var updatedClient model.OidcClient
err = db.First(&updatedClient, "id = ?", client.ID).Error
require.NoError(t, err)
require.NotNil(t, updatedClient.ImageType)
assert.Equal(t, "webp", *updatedClient.ImageType)
// Old PNG file should be deleted
assert.False(t, fileExists(t, oldPath), "Old PNG file should have been deleted")
})
t.Run("Updates dark logo type and deletes old file when type changes", func(t *testing.T) {
// Create the old JPG dark file in storage
oldPath := "oidc-client-images/" + client.ID + "-dark.jpg"
createDummyFile(t, oldPath)
require.True(t, fileExists(t, oldPath), "Old dark file should exist before update")
// Client currently has a JPG dark logo, update to WEBP
err := s.updateClientLogoType(t.Context(), client.ID, "webp", false)
require.NoError(t, err)
// Verify the client was updated
var updatedClient model.OidcClient
err = db.First(&updatedClient, "id = ?", client.ID).Error
require.NoError(t, err)
require.NotNil(t, updatedClient.DarkImageType)
assert.Equal(t, "webp", *updatedClient.DarkImageType)
// Old JPG dark file should be deleted
assert.False(t, fileExists(t, oldPath), "Old JPG dark file should have been deleted")
})
t.Run("Does not delete file when type remains the same", func(t *testing.T) {
// Create the WEBP file in storage
webpPath := "oidc-client-images/" + client.ID + ".webp"
createDummyFile(t, webpPath)
require.True(t, fileExists(t, webpPath), "WEBP file should exist before update")
// Update to the same type (WEBP)
err := s.updateClientLogoType(t.Context(), client.ID, "webp", true)
require.NoError(t, err)
// Verify the client still has WEBP
var updatedClient model.OidcClient
err = db.First(&updatedClient, "id = ?", client.ID).Error
require.NoError(t, err)
require.NotNil(t, updatedClient.ImageType)
assert.Equal(t, "webp", *updatedClient.ImageType)
// WEBP file should still exist since type didn't change
assert.True(t, fileExists(t, webpPath), "WEBP file should still exist")
})
t.Run("Returns error for non-existent client", func(t *testing.T) {
err := s.updateClientLogoType(t.Context(), "non-existent-client-id", "png", true)
require.Error(t, err)
require.True(t, apperror.IsCode(err, apperror.CodeNotFound))
})
}
func TestOidcClientImagePath(t *testing.T) {
const metadataClientID = "https://app.example.com/oauth/client"
assert.Equal(t, "oidc-client-images/client-id.png", oidcClientImagePath("client-id", "", "png"))
assert.Equal(
t,
"oidc-client-images/cimd-"+utils.CreateSha256Hash(metadataClientID)+"-dark.webp",
oidcClientImagePath(metadataClientID, "-dark", "webp"),
)
assert.NotContains(t, oidcClientImagePath(metadataClientID, "", "png"), "app.example.com")
}
func TestOidcService_downloadAndSaveLogoFromURL(t *testing.T) {
const publicLogoHost = "https://8.8.8.8"
// Create a test database
db := testutils.NewDatabaseForTest(t)
// Create database storage
dbStorage, err := storage.NewDatabaseStorage(db)
require.NoError(t, err)
// Create a test client
client := model.OidcClient{
Name: "Test Client",
CallbackURLs: datatype.StringList{"https://example.com/callback"},
}
err = db.Create(&client).Error
require.NoError(t, err)
// Helper function to check if a file exists in storage
fileExists := func(t *testing.T, path string) bool {
t.Helper()
_, _, err := dbStorage.Open(t.Context(), path)
return err == nil
}
// Helper function to get file content from storage
getFileContent := func(t *testing.T, path string) []byte {
t.Helper()
reader, _, err := dbStorage.Open(t.Context(), path)
require.NoError(t, err)
defer reader.Close()
content, err := io.ReadAll(reader)
require.NoError(t, err)
return content
}
t.Run("Successfully downloads and saves PNG logo from URL", func(t *testing.T) {
// Create mock PNG content
pngContent := []byte("fake-png-content")
// Create a mock HTTP response with headers
//nolint:bodyclose
pngResponse := testutils.NewMockResponse(http.StatusOK, string(pngContent))
pngResponse.Header.Set("Content-Type", "image/png")
// Create a mock HTTP client with responses
mockResponses := map[string]*http.Response{
//nolint:bodyclose
publicLogoHost + "/logo.png": pngResponse,
}
httpClient := &http.Client{
Transport: &testutils.MockRoundTripper{
Responses: mockResponses,
},
}
// Init the OidcService with mock HTTP client
s := &OidcService{
db: db,
fileStorage: dbStorage,
httpClient: httpClient,
}
// Download and save the logo
err := s.downloadAndSaveLogoFromURL(t.Context(), client.ID, publicLogoHost+"/logo.png", true)
require.NoError(t, err)
// Verify the file was saved
logoPath := "oidc-client-images/" + client.ID + ".png"
require.True(t, fileExists(t, logoPath), "Logo file should exist in storage")
// Verify the content
savedContent := getFileContent(t, logoPath)
assert.Equal(t, pngContent, savedContent)
// Verify the client was updated
var updatedClient model.OidcClient
err = db.First(&updatedClient, "id = ?", client.ID).Error
require.NoError(t, err)
require.NotNil(t, updatedClient.ImageType)
assert.Equal(t, "png", *updatedClient.ImageType)
})
t.Run("Successfully downloads and saves dark logo", func(t *testing.T) {
// Create mock WEBP content
webpContent := []byte("fake-webp-content")
//nolint:bodyclose
webpResponse := testutils.NewMockResponse(http.StatusOK, string(webpContent))
webpResponse.Header.Set("Content-Type", "image/webp")
mockResponses := map[string]*http.Response{
//nolint:bodyclose
publicLogoHost + "/dark-logo.webp": webpResponse,
}
httpClient := &http.Client{
Transport: &testutils.MockRoundTripper{
Responses: mockResponses,
},
}
s := &OidcService{
db: db,
fileStorage: dbStorage,
httpClient: httpClient,
}
// Download and save the dark logo
err := s.downloadAndSaveLogoFromURL(t.Context(), client.ID, publicLogoHost+"/dark-logo.webp", false)
require.NoError(t, err)
// Verify the dark logo file was saved
darkLogoPath := "oidc-client-images/" + client.ID + "-dark.webp"
require.True(t, fileExists(t, darkLogoPath), "Dark logo file should exist in storage")
// Verify the content
savedContent := getFileContent(t, darkLogoPath)
assert.Equal(t, webpContent, savedContent)
// Verify the client was updated
var updatedClient model.OidcClient
err = db.First(&updatedClient, "id = ?", client.ID).Error
require.NoError(t, err)
require.NotNil(t, updatedClient.DarkImageType)
assert.Equal(t, "webp", *updatedClient.DarkImageType)
})
t.Run("Detects extension from URL path", func(t *testing.T) {
svgContent := []byte("<svg></svg>")
mockResponses := map[string]*http.Response{
//nolint:bodyclose
publicLogoHost + "/icon.svg": testutils.NewMockResponse(http.StatusOK, string(svgContent)),
}
httpClient := &http.Client{
Transport: &testutils.MockRoundTripper{
Responses: mockResponses,
},
}
s := &OidcService{
db: db,
fileStorage: dbStorage,
httpClient: httpClient,
}
err := s.downloadAndSaveLogoFromURL(t.Context(), client.ID, publicLogoHost+"/icon.svg", true)
require.NoError(t, err)
// Verify SVG file was saved
logoPath := "oidc-client-images/" + client.ID + ".svg"
require.True(t, fileExists(t, logoPath), "SVG logo should exist")
})
t.Run("Detects extension from Content-Type when path has no extension", func(t *testing.T) {
jpgContent := []byte("fake-jpg-content")
//nolint:bodyclose
jpgResponse := testutils.NewMockResponse(http.StatusOK, string(jpgContent))
jpgResponse.Header.Set("Content-Type", "image/jpeg")
mockResponses := map[string]*http.Response{
//nolint:bodyclose
publicLogoHost + "/logo": jpgResponse,
}
httpClient := &http.Client{
Transport: &testutils.MockRoundTripper{
Responses: mockResponses,
},
}
s := &OidcService{
db: db,
fileStorage: dbStorage,
httpClient: httpClient,
}
err := s.downloadAndSaveLogoFromURL(t.Context(), client.ID, publicLogoHost+"/logo", true)
require.NoError(t, err)
// Verify JPG file was saved (jpeg extension is normalized to jpg)
logoPath := "oidc-client-images/" + client.ID + ".jpg"
require.True(t, fileExists(t, logoPath), "JPG logo should exist")
})
t.Run("Returns error for invalid URL", func(t *testing.T) {
s := &OidcService{
db: db,
fileStorage: dbStorage,
httpClient: &http.Client{},
}
err := s.downloadAndSaveLogoFromURL(t.Context(), client.ID, "://invalid-url", true)
require.Error(t, err)
require.True(t, apperror.IsCode(err, apperror.CodeValidationFailed))
})
t.Run("Returns error for non-200 status code", func(t *testing.T) {
mockResponses := map[string]*http.Response{
//nolint:bodyclose
publicLogoHost + "/not-found.png": testutils.NewMockResponse(http.StatusNotFound, "Not Found"),
}
httpClient := &http.Client{
Transport: &testutils.MockRoundTripper{
Responses: mockResponses,
},
}
s := &OidcService{
db: db,
fileStorage: dbStorage,
httpClient: httpClient,
}
err := s.downloadAndSaveLogoFromURL(t.Context(), client.ID, publicLogoHost+"/not-found.png", true)
require.Error(t, err)
require.True(t, apperror.IsCode(err, apperror.CodeLogoDownloadFailed))
})
t.Run("Returns error for too large content", func(t *testing.T) {
// Create content larger than 2MB (maxLogoSize)
largeContent := strings.Repeat("x", 2<<20+100) // 2.1MB
//nolint:bodyclose
largeResponse := testutils.NewMockResponse(http.StatusOK, largeContent)
largeResponse.Header.Set("Content-Type", "image/png")
largeResponse.Header.Set("Content-Length", strconv.Itoa(len(largeContent)))
mockResponses := map[string]*http.Response{
//nolint:bodyclose
publicLogoHost + "/large.png": largeResponse,
}
httpClient := &http.Client{
Transport: &testutils.MockRoundTripper{
Responses: mockResponses,
},
}
s := &OidcService{
db: db,
fileStorage: dbStorage,
httpClient: httpClient,
}
err := s.downloadAndSaveLogoFromURL(t.Context(), client.ID, publicLogoHost+"/large.png", true)
require.Error(t, err)
require.True(t, apperror.IsCode(err, apperror.CodeLogoTooLarge))
})
t.Run("Returns error for unsupported file type", func(t *testing.T) {
//nolint:bodyclose
textResponse := testutils.NewMockResponse(http.StatusOK, "text content")
textResponse.Header.Set("Content-Type", "text/plain")
mockResponses := map[string]*http.Response{
//nolint:bodyclose
publicLogoHost + "/file.txt": textResponse,
}
httpClient := &http.Client{
Transport: &testutils.MockRoundTripper{
Responses: mockResponses,
},
}
s := &OidcService{
db: db,
fileStorage: dbStorage,
httpClient: httpClient,
}
err := s.downloadAndSaveLogoFromURL(t.Context(), client.ID, publicLogoHost+"/file.txt", true)
require.Error(t, err)
require.True(t, apperror.IsCode(err, apperror.CodeLogoTypeNotSupported))
})
t.Run("Returns error for non-existent client", func(t *testing.T) {
//nolint:bodyclose
pngResponse := testutils.NewMockResponse(http.StatusOK, "content")
pngResponse.Header.Set("Content-Type", "image/png")
mockResponses := map[string]*http.Response{
//nolint:bodyclose
publicLogoHost + "/logo.png": pngResponse,
}
httpClient := &http.Client{
Transport: &testutils.MockRoundTripper{
Responses: mockResponses,
},
}
s := &OidcService{
db: db,
fileStorage: dbStorage,
httpClient: httpClient,
}
err := s.downloadAndSaveLogoFromURL(t.Context(), "non-existent-client-id", publicLogoHost+"/logo.png", true)
require.Error(t, err)
require.True(t, apperror.IsCode(err, apperror.CodeNotFound))
})
}
func TestOidcService_CreateClient_withDescription(t *testing.T) {
db := testutils.NewDatabaseForTest(t)
s, err := NewOidcService(db, nil, nil, nil, nil, nil, nil)
require.NoError(t, err)
description := "A test client description"
input := dto.OidcClientCreateDto{
OidcClientUpdateDto: dto.OidcClientUpdateDto{
Name: "Test Client",
Description: description,
CallbackURLs: []string{"https://example.com/callback"},
},
}
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)
require.NotEmpty(t, fetched.Description)
assert.Equal(t, description, fetched.Description)
}
func TestOidcService_CreateClient_withoutDescription(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"},
},
}
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.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)
s, err := NewOidcService(db, nil, nil, nil, nil, nil, nil)
require.NoError(t, err)
client := model.OidcClient{Name: "Test Client"}
err = db.Create(&client).Error
require.NoError(t, err)
customSecret := "custom-client-secret-with-a-minimum-length"
input := dto.OidcClientSecretCreateDto{Secret: customSecret}
created, secret, err := s.CreateClientSecret(t.Context(), client.ID, input)
require.NoError(t, err)
assert.Equal(t, customSecret, secret)
assert.Equal(t, "cust", created.Prefix)
assert.Nil(t, created.ExpiresAt)
assert.True(t, created.IsActive())
var fetched model.OidcClient
err = db.First(&fetched, "id = ?", client.ID).Error
require.NoError(t, err)
require.Len(t, fetched.Credentials.Secrets, 1)
assert.Equal(t, created.ID, fetched.Credentials.Secrets[0].ID)
assert.Equal(t, model.OidcClientSecretHashSHA256, fetched.Credentials.Secrets[0].Algorithm)
assert.Equal(t, utils.CreateSha256Hash(customSecret), fetched.Credentials.Secrets[0].Hash)
}
func TestOidcService_CreateClientSecret_multipleSecrets(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"}
err = db.Create(&client).Error
require.NoError(t, err)
// Adding a second secret leaves the first one in place
first, _, err := s.CreateClientSecret(t.Context(), client.ID, dto.OidcClientSecretCreateDto{})
require.NoError(t, err)
expiresAt := datatype.DateTime(time.Now().Add(24 * time.Hour))
second, _, err := s.CreateClientSecret(t.Context(), client.ID, dto.OidcClientSecretCreateDto{ExpiresAt: &expiresAt})
require.NoError(t, err)
secrets, err := s.ListClientSecrets(t.Context(), client.ID)
require.NoError(t, err)
require.Len(t, secrets, 2)
assert.Equal(t, first.ID, secrets[0].ID)
assert.Equal(t, second.ID, secrets[1].ID)
assert.Equal(t, expiresAt.ToTime().Unix(), secrets[1].ExpiresAt.ToTime().Unix())
// Deleting one secret keeps the other usable
err = s.DeleteClientSecret(t.Context(), client.ID, first.ID)
require.NoError(t, err)
secrets, err = s.ListClientSecrets(t.Context(), client.ID)
require.NoError(t, err)
require.Len(t, secrets, 1)
assert.Equal(t, second.ID, secrets[0].ID)
// Deleting an unknown secret is reported as not found
err = s.DeleteClientSecret(t.Context(), client.ID, first.ID)
require.Error(t, err)
}
func TestOidcService_CreateClientSecret_expirationInThePast(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"}
err = db.Create(&client).Error
require.NoError(t, err)
expiresAt := datatype.DateTime(time.Now().Add(-time.Minute))
_, _, err = s.CreateClientSecret(t.Context(), client.ID, dto.OidcClientSecretCreateDto{ExpiresAt: &expiresAt})
require.Error(t, err)
}
func TestOidcService_CreateClientSecret_limit(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"}
err = db.Create(&client).Error
require.NoError(t, err)
for range model.MaxOidcClientSecrets {
_, _, err = s.CreateClientSecret(t.Context(), client.ID, dto.OidcClientSecretCreateDto{})
require.NoError(t, err)
}
_, _, err = s.CreateClientSecret(t.Context(), client.ID, dto.OidcClientSecretCreateDto{})
require.Error(t, err)
}
func TestOidcService_CreateClientSecret_preservesFederatedIdentities(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"},
Credentials: model.OidcClientCredentials{
FederatedIdentities: []model.OidcClientFederatedIdentity{{Issuer: "https://issuer.example.com"}},
},
}
err = db.Create(&client).Error
require.NoError(t, err)
_, _, err = s.CreateClientSecret(t.Context(), client.ID, dto.OidcClientSecretCreateDto{})
require.NoError(t, err)
// Updating the client must not drop the secrets managed by the dedicated endpoints
_, err = s.UpdateClient(t.Context(), client.ID, dto.OidcClientUpdateDto{
Name: "Test Client",
CallbackURLs: []string{"https://example.com/callback"},
Credentials: dto.OidcClientCredentialsDto{
FederatedIdentities: []dto.OidcClientFederatedIdentityDto{{Issuer: "https://other.example.com"}},
},
})
require.NoError(t, err)
var fetched model.OidcClient
require.NoError(t, db.First(&fetched, "id = ?", client.ID).Error)
require.Len(t, fetched.Credentials.Secrets, 1)
require.Len(t, fetched.Credentials.FederatedIdentities, 1)
assert.Equal(t, "https://other.example.com", fetched.Credentials.FederatedIdentities[0].Issuer)
}
func TestOidcService_UpdateClient_description(t *testing.T) {
db := testutils.NewDatabaseForTest(t)
s, err := NewOidcService(db, nil, nil, nil, nil, nil, nil)
require.NoError(t, err)
// Create a client without a description
client := model.OidcClient{
Name: "Test Client",
CallbackURLs: datatype.StringList{"https://example.com/callback"},
}
err = db.Create(&client).Error
require.NoError(t, err)
// Update with a description
description := "Updated description"
input := dto.OidcClientUpdateDto{
Name: "Test Client",
Description: description,
CallbackURLs: []string{"https://example.com/callback"},
}
_, err = s.UpdateClient(t.Context(), client.ID, input)
require.NoError(t, err)
var fetched model.OidcClient
err = db.First(&fetched, "id = ?", client.ID).Error
require.NoError(t, err)
require.NotEmpty(t, fetched.Description)
assert.Equal(t, description, fetched.Description)
// Update to clear the description
input.Description = ""
_, err = s.UpdateClient(t.Context(), client.ID, input)
require.NoError(t, err)
err = db.First(&fetched, "id = ?", client.ID).Error
require.NoError(t, err)
assert.Empty(t, fetched.Description)
}
func TestOidcService_UpdateClient_CIMDPreservesMetadataFields(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: "Metadata Client",
CallbackURLs: datatype.StringList{"https://metadata.example.com/callback"},
LogoutCallbackURLs: datatype.StringList{"https://metadata.example.com/logout"},
IsPublic: true,
PkceEnabled: true,
Credentials: model.OidcClientCredentials{
FederatedIdentities: []model.OidcClientFederatedIdentity{{
Issuer: "https://metadata.example.com/client.json",
Subject: "https://metadata.example.com/client.json",
JWKS: "https://metadata.example.com/jwks.json",
}},
},
ClientType: model.OidcClientTypeCIMD,
}
require.NoError(t, db.Create(&client).Error)
launchURL := "https://app.example.com"
accessDuration := int64(2 * 60)
refreshDuration := int64(7 * 24 * 60)
input := dto.OidcClientUpdateDto{
Name: "Overridden Client",
Description: "Locally managed description",
CallbackURLs: []string{"https://override.example.com/callback"},
LogoutCallbackURLs: []string{"https://override.example.com/logout"},
IsPublic: false,
PkceEnabled: false,
RequiresReauthentication: true,
RequiresPushedAuthorizationRequests: true,
SkipConsent: true,
LaunchURL: &launchURL,
IsGroupRestricted: true,
AccessTokenDurationMinutes: accessDuration,
RefreshTokenDurationMinutes: refreshDuration,
Credentials: dto.OidcClientCredentialsDto{
FederatedIdentities: []dto.OidcClientFederatedIdentityDto{{
Issuer: "https://override.example.com",
JWKS: "https://override.example.com/jwks.json",
}},
},
}
_, 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, client.Name, fetched.Name)
assert.Equal(t, client.CallbackURLs, fetched.CallbackURLs)
assert.Equal(t, client.LogoutCallbackURLs, fetched.LogoutCallbackURLs)
assert.Equal(t, client.IsPublic, fetched.IsPublic)
assert.Equal(t, client.PkceEnabled, fetched.PkceEnabled)
assert.Equal(t, client.Credentials, fetched.Credentials)
assert.Equal(t, input.Description, fetched.Description)
assert.Equal(t, input.RequiresReauthentication, fetched.RequiresReauthentication)
assert.Equal(t, input.RequiresPushedAuthorizationRequests, fetched.RequiresPushedAuthorizationRequests)
assert.Equal(t, input.SkipConsent, fetched.SkipConsent)
assert.Equal(t, input.LaunchURL, fetched.LaunchURL)
assert.Equal(t, input.IsGroupRestricted, fetched.IsGroupRestricted)
assert.Equal(t, accessDuration, fetched.AccessTokenDurationMinutes)
assert.Equal(t, refreshDuration, fetched.RefreshTokenDurationMinutes)
}
func TestOidcService_UpdateClient_CIMDDoesNotOverwriteConcurrentMetadataRefresh(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: "Original metadata name",
CallbackURLs: datatype.StringList{"https://metadata.example.com/callback"},
ClientType: model.OidcClientTypeCIMD,
}
require.NoError(t, db.Create(&client).Error)
// Simulate metadata refresh changing a document-owned column after the admin request read its snapshot
require.NoError(t, db.Exec(`
CREATE TRIGGER refresh_metadata_before_admin_update
BEFORE UPDATE OF description ON oidc_clients
BEGIN
UPDATE oidc_clients SET name = 'Refreshed metadata name' WHERE id = OLD.id;
END;
`).Error)
input := dto.OidcClientUpdateDto{
Description: "Locally managed description",
}
_, 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, "Refreshed metadata name", fetched.Name)
assert.Equal(t, input.Description, fetched.Description)
}
func TestOidcService_ListAccessibleOidcClients_requiresExplicitGroupPermission(t *testing.T) {
db := testutils.NewDatabaseForTest(t)
s, err := NewOidcService(db, nil, nil, nil, nil, nil, nil)
require.NoError(t, err)
allowedGroup := model.UserGroup{Name: "allowed", FriendlyName: "Allowed"}
otherGroup := model.UserGroup{Name: "other", FriendlyName: "Other"}
require.NoError(t, db.Create(&allowedGroup).Error)
require.NoError(t, db.Create(&otherGroup).Error)
userWithGroup := model.User{Username: "with-group", UserGroups: []model.UserGroup{allowedGroup}}
userWithoutGroup := model.User{Username: "without-group"}
require.NoError(t, db.Create(&userWithGroup).Error)
require.NoError(t, db.Create(&userWithoutGroup).Error)
clients := []model.OidcClient{
{Name: "Unrestricted", CallbackURLs: datatype.StringList{"https://unrestricted.example.com/callback"}},
{Name: "Restricted without groups", CallbackURLs: datatype.StringList{"https://empty.example.com/callback"}, IsGroupRestricted: true},
{Name: "Restricted to user group", CallbackURLs: datatype.StringList{"https://allowed.example.com/callback"}, IsGroupRestricted: true, AllowedUserGroups: []model.UserGroup{allowedGroup}},
{Name: "Restricted to other group", CallbackURLs: datatype.StringList{"https://other.example.com/callback"}, IsGroupRestricted: true, AllowedUserGroups: []model.UserGroup{otherGroup}},
}
for i := range clients {
require.NoError(t, db.Create(&clients[i]).Error)
}
groupClients, _, err := s.ListAccessibleOidcClients(t.Context(), userWithGroup.ID, utils.ListRequestOptions{})
require.NoError(t, err)
assert.ElementsMatch(t, []string{"Unrestricted", "Restricted to user group"}, accessibleClientNames(groupClients))
noGroupClients, _, err := s.ListAccessibleOidcClients(t.Context(), userWithoutGroup.ID, utils.ListRequestOptions{})
require.NoError(t, err)
assert.Equal(t, []string{"Unrestricted"}, accessibleClientNames(noGroupClients))
}
func TestOidcService_ListClientViewsFilterByLaunchURLPresence(t *testing.T) {
db := testutils.NewDatabaseForTest(t)
s, err := NewOidcService(db, nil, nil, nil, nil, nil, nil)
require.NoError(t, err)
user := model.User{Username: "launch-url-filter"}
require.NoError(t, db.Create(&user).Error)
launchURL := "https://launchable.example.com"
emptyLaunchURL := ""
clients := []model.OidcClient{
{Name: "Launchable", LaunchURL: &launchURL},
{Name: "Missing launch URL"},
{Name: "Empty launch URL", LaunchURL: &emptyLaunchURL},
}
for i := range clients {
require.NoError(t, db.Create(&clients[i]).Error)
require.NoError(t, db.Create(&model.UserAuthorizedOidcClient{
UserID: user.ID,
ClientID: clients[i].ID,
}).Error)
}
withLaunchURL := utils.ListRequestOptions{
Filters: map[string][]any{"hasLaunchURL": {true}},
}
withoutLaunchURL := utils.ListRequestOptions{
Filters: map[string][]any{"hasLaunchURL": {false}},
}
allClients, allClientsPagination, err := s.ListAccessibleOidcClients(t.Context(), user.ID, utils.ListRequestOptions{})
require.NoError(t, err)
assert.Equal(t, int64(3), allClientsPagination.TotalItems)
assert.ElementsMatch(t, []string{"Launchable", "Missing launch URL", "Empty launch URL"}, accessibleClientNames(allClients))
launchableClients, launchablePagination, err := s.ListAccessibleOidcClients(t.Context(), user.ID, withLaunchURL)
require.NoError(t, err)
assert.Equal(t, int64(1), launchablePagination.TotalItems)
assert.Equal(t, []string{"Launchable"}, accessibleClientNames(launchableClients))
allAuthorizations, allAuthorizationsPagination, err := s.ListAuthorizedClients(t.Context(), user.ID, utils.ListRequestOptions{})
require.NoError(t, err)
assert.Equal(t, int64(3), allAuthorizationsPagination.TotalItems)
assert.Len(t, allAuthorizations, 3)
hiddenAuthorizations, hiddenPagination, err := s.ListAuthorizedClients(t.Context(), user.ID, withoutLaunchURL)
require.NoError(t, err)
assert.Equal(t, int64(2), hiddenPagination.TotalItems)
assert.ElementsMatch(t, []string{"Missing launch URL", "Empty launch URL"}, []string{
hiddenAuthorizations[0].Client.Name,
hiddenAuthorizations[1].Client.Name,
})
}
func accessibleClientNames(clients []dto.AccessibleOidcClientDto) []string {
names := make([]string, len(clients))
for i := range clients {
names[i] = clients[i].Name
}
return names
}