package service import ( "context" "crypto/rand" "crypto/rsa" "crypto/x509" "crypto/x509/pkix" "encoding/pem" "math/big" "testing" "time" "github.com/fleetdm/fleet/v4/server/authz" "github.com/fleetdm/fleet/v4/server/config" "github.com/fleetdm/fleet/v4/server/fleet" "github.com/fleetdm/fleet/v4/server/mock" "github.com/fleetdm/fleet/v4/server/test" "github.com/jmoiron/sqlx" "github.com/stretchr/testify/require" ) func TestConditionalAccessGetIdPSigningCertAuth(t *testing.T) { t.Parallel() ds := new(mock.Store) cfg := config.TestConfig() cfg.Server.PrivateKey = "test-private-key" svc, ctx := newTestServiceWithConfig(t, ds, cfg, nil, nil) // Mock the datastore to return a valid IdP certificate ds.GetAllMDMConfigAssetsByNameFunc = func(ctx context.Context, assetNames []fleet.MDMAssetName, _ sqlx.QueryerContext) (map[fleet.MDMAssetName]fleet.MDMConfigAsset, error) { return map[fleet.MDMAssetName]fleet.MDMConfigAsset{ fleet.MDMAssetConditionalAccessIDPCert: { Name: fleet.MDMAssetConditionalAccessIDPCert, Value: []byte("-----BEGIN CERTIFICATE-----\ntest\n-----END CERTIFICATE-----"), }, }, nil } testCases := []struct { name string user *fleet.User shouldFail bool }{ {"global admin", test.UserAdmin, false}, {"global maintainer", test.UserMaintainer, false}, {"global observer", test.UserObserver, false}, {"global observer+", test.UserObserverPlus, false}, {"global gitops", test.UserGitOps, false}, {"team admin", test.UserTeamAdminTeam1, true}, {"team maintainer", test.UserTeamMaintainerTeam1, true}, {"team observer", test.UserTeamObserverTeam1, true}, {"team observer+", test.UserTeamObserverPlusTeam1, true}, {"team gitops", test.UserTeamGitOpsTeam1, true}, {"user no roles", test.UserNoRoles, true}, } for _, tt := range testCases { t.Run(tt.name, func(t *testing.T) { ctx := test.UserContext(ctx, tt.user) certPEM, err := svc.ConditionalAccessGetIdPSigningCert(ctx) if tt.shouldFail { require.Error(t, err) var forbiddenError *authz.Forbidden require.ErrorAs(t, err, &forbiddenError) require.Nil(t, certPEM) } else { require.NoError(t, err) require.NotNil(t, certPEM) } }) } } func TestConditionalAccessGetIdPSigningCert(t *testing.T) { t.Parallel() t.Run("missing server private key", func(t *testing.T) { ds := new(mock.Store) cfg := config.TestConfig() cfg.Server.PrivateKey = "" // Not configured svc, ctx := newTestServiceWithConfig(t, ds, cfg, nil, nil) ctx = test.UserContext(ctx, test.UserAdmin) certPEM, err := svc.ConditionalAccessGetIdPSigningCert(ctx) var badReqErr *fleet.BadRequestError require.ErrorAs(t, err, &badReqErr) require.Contains(t, err.Error(), "Fleet server private key is not configured") require.Nil(t, certPEM) }) } func TestConditionalAccessGetIdPAppleProfileAuth(t *testing.T) { t.Parallel() ds := new(mock.Store) cfg := config.TestConfig() cfg.Server.PrivateKey = "test-private-key" svc, ctx := newTestServiceWithConfig(t, ds, cfg, nil, nil) // Mock valid certificate certPEM := generateTestCertPEM(t) // Mock the datastore methods ds.GetAllMDMConfigAssetsByNameFunc = func(ctx context.Context, assetNames []fleet.MDMAssetName, _ sqlx.QueryerContext) (map[fleet.MDMAssetName]fleet.MDMConfigAsset, error) { return map[fleet.MDMAssetName]fleet.MDMConfigAsset{ fleet.MDMAssetConditionalAccessCACert: { Name: fleet.MDMAssetConditionalAccessCACert, Value: certPEM, }, }, nil } ds.AppConfigFunc = func(ctx context.Context) (*fleet.AppConfig, error) { return &fleet.AppConfig{ ServerSettings: fleet.ServerSettings{ ServerURL: "https://fleet.example.com", }, }, nil } ds.GetEnrollSecretsFunc = func(ctx context.Context, teamID *uint) ([]*fleet.EnrollSecret, error) { return []*fleet.EnrollSecret{ {Secret: "test-secret-123"}, }, nil } testCases := []struct { name string user *fleet.User shouldFail bool }{ {"global admin", test.UserAdmin, false}, {"global maintainer", test.UserMaintainer, false}, {"global observer", test.UserObserver, false}, {"global observer+", test.UserObserverPlus, false}, {"global gitops", test.UserGitOps, false}, {"team admin", test.UserTeamAdminTeam1, true}, {"team maintainer", test.UserTeamMaintainerTeam1, true}, {"team observer", test.UserTeamObserverTeam1, true}, {"team observer+", test.UserTeamObserverPlusTeam1, true}, {"team gitops", test.UserTeamGitOpsTeam1, true}, {"user no roles", test.UserNoRoles, true}, } for _, tt := range testCases { t.Run(tt.name, func(t *testing.T) { ctx := test.UserContext(ctx, tt.user) profileData, err := svc.ConditionalAccessGetIdPAppleProfile(ctx) if tt.shouldFail { require.Error(t, err) var forbiddenError *authz.Forbidden require.ErrorAs(t, err, &forbiddenError) require.Nil(t, profileData) } else { require.NoError(t, err) require.NotNil(t, profileData) // Verify the profile contains expected content profileStr := string(profileData) require.Contains(t, profileStr, "com.fleetdm.conditional-access") require.Contains(t, profileStr, "https://okta.fleet.example.com") } }) } } // generateTestCertPEM generates a test certificate in PEM format for testing func generateTestCertPEM(t *testing.T) []byte { // Create a simple self-signed certificate template := &x509.Certificate{ SerialNumber: big.NewInt(1), Subject: pkix.Name{ CommonName: "Test CA", }, NotBefore: time.Now(), NotAfter: time.Now().Add(24 * time.Hour), IsCA: true, BasicConstraintsValid: true, } priv, err := rsa.GenerateKey(rand.Reader, 2048) require.NoError(t, err) certBytes, err := x509.CreateCertificate(rand.Reader, template, template, &priv.PublicKey, priv) require.NoError(t, err) certPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: certBytes}) return certPEM } func TestConditionalAccessGetIdPAppleProfile(t *testing.T) { certPEM := generateTestCertPEM(t) t.Run("missing server private key", func(t *testing.T) { ds := new(mock.Store) cfg := config.TestConfig() cfg.Server.PrivateKey = "" // Not configured svc, ctx := newTestServiceWithConfig(t, ds, cfg, nil, nil) ctx = test.UserContext(ctx, test.UserAdmin) profileData, err := svc.ConditionalAccessGetIdPAppleProfile(ctx) var badReqErr *fleet.BadRequestError require.ErrorAs(t, err, &badReqErr) require.Contains(t, err.Error(), "Fleet server private key is not configured") require.Nil(t, profileData) }) t.Run("success - generates valid profile", func(t *testing.T) { ds := new(mock.Store) cfg := config.TestConfig() cfg.Server.PrivateKey = "test-private-key" svc, ctx := newTestServiceWithConfig(t, ds, cfg, nil, nil) ctx = test.UserContext(ctx, test.UserAdmin) ds.GetAllMDMConfigAssetsByNameFunc = func(ctx context.Context, assetNames []fleet.MDMAssetName, _ sqlx.QueryerContext) (map[fleet.MDMAssetName]fleet.MDMConfigAsset, error) { return map[fleet.MDMAssetName]fleet.MDMConfigAsset{ fleet.MDMAssetConditionalAccessCACert: { Name: fleet.MDMAssetConditionalAccessCACert, Value: certPEM, }, }, nil } ds.AppConfigFunc = func(ctx context.Context) (*fleet.AppConfig, error) { return &fleet.AppConfig{ ServerSettings: fleet.ServerSettings{ ServerURL: "https://fleet.example.com:8080", }, }, nil } ds.GetEnrollSecretsFunc = func(ctx context.Context, teamID *uint) ([]*fleet.EnrollSecret, error) { return []*fleet.EnrollSecret{ {Secret: "test-secret-456"}, }, nil } profileData, err := svc.ConditionalAccessGetIdPAppleProfile(ctx) require.NoError(t, err) require.NotEmpty(t, profileData) profileStr := string(profileData) // Verify XML structure require.Contains(t, profileStr, "") require.Contains(t, profileStr, "