diff --git a/server/datastore/mysql/android.go b/server/datastore/mysql/android.go index 27228ac21d..dfe8d57407 100644 --- a/server/datastore/mysql/android.go +++ b/server/datastore/mysql/android.go @@ -15,10 +15,6 @@ import ( "github.com/jmoiron/sqlx" ) -func (ds *Datastore) GetAndroidDS() android.Datastore { - return ds.androidDS -} - func (ds *Datastore) NewAndroidHost(ctx context.Context, host *fleet.AndroidHost) (*fleet.AndroidHost, error) { if !host.IsValid() { return nil, ctxerr.New(ctx, "valid Android host is required") @@ -116,7 +112,7 @@ func (ds *Datastore) NewAndroidHost(ctx context.Context, host *fleet.AndroidHost return ctxerr.Wrap(ctx, err, "new Android host MDM info") } - host.Device, err = ds.androidDS.CreateDeviceTx(ctx, tx, host.Device) + host.Device, err = ds.Datastore.CreateDeviceTx(ctx, tx, host.Device) if err != nil { return ctxerr.Wrap(ctx, err, "creating new Android device") } @@ -191,7 +187,7 @@ func (ds *Datastore) UpdateAndroidHost(ctx context.Context, host *fleet.AndroidH } } - err = ds.androidDS.UpdateDeviceTx(ctx, tx, host.Device) + err = ds.Datastore.UpdateDeviceTx(ctx, tx, host.Device) if err != nil { return ctxerr.Wrap(ctx, err, "update Android device") } diff --git a/server/datastore/mysql/android_test.go b/server/datastore/mysql/android_test.go index 5d427a8533..b322a47b8e 100644 --- a/server/datastore/mysql/android_test.go +++ b/server/datastore/mysql/android_test.go @@ -243,7 +243,7 @@ func testAndroidMDMStats(t *testing.T, ds *Datastore) { require.Equal(t, serverURL, solutionsStats[0].ServerURL) // turn MDM off for android - err = ds.androidDS.DeleteAllEnterprises(testCtx()) + err = ds.DeleteAllEnterprises(testCtx()) require.NoError(t, err) err = ds.BulkSetAndroidHostsUnenrolled(testCtx()) require.NoError(t, err) diff --git a/server/datastore/mysql/mysql.go b/server/datastore/mysql/mysql.go index bc9c0a2643..80063e4423 100644 --- a/server/datastore/mysql/mysql.go +++ b/server/datastore/mysql/mysql.go @@ -51,11 +51,11 @@ type Datastore struct { replica fleet.DBReader // so it cannot be used to perform writes primary *sqlx.DB - logger log.Logger - clock clock.Clock - config config.MysqlConfig - pusher nano_push.Pusher - androidDS android.Datastore + logger log.Logger + clock clock.Clock + config config.MysqlConfig + pusher nano_push.Pusher + android.Datastore // nil if no read replica readReplicaConfig *config.MysqlConfig @@ -259,7 +259,7 @@ func New(config config.MysqlConfig, c clock.Clock, opts ...DBOption) (*Datastore stmtCache: make(map[string]*sqlx.Stmt), minLastOpenedAtDiff: options.MinLastOpenedAtDiff, serverPrivateKey: options.PrivateKey, - androidDS: android_mysql.New(options.Logger, dbWriter, dbReader), + Datastore: android_mysql.New(options.Logger, dbWriter, dbReader), } go ds.writeChanLoop() diff --git a/server/fleet/datastore.go b/server/fleet/datastore.go index c6989d4109..5c1c7d92a9 100644 --- a/server/fleet/datastore.go +++ b/server/fleet/datastore.go @@ -394,14 +394,6 @@ type Datastore interface { OSVersion(ctx context.Context, osVersionID uint, teamFilter *TeamFilter) (*OSVersion, *time.Time, error) UpdateOSVersions(ctx context.Context) error - // //////////////////////////////////////////////////////////////////////////// - // Android - - GetAndroidDS() android.Datastore - NewAndroidHost(ctx context.Context, host *AndroidHost) (*AndroidHost, error) - UpdateAndroidHost(ctx context.Context, host *AndroidHost, fromEnroll bool) error - AndroidHostLite(ctx context.Context, enterpriseSpecificID string) (*AndroidHost, error) - /////////////////////////////////////////////////////////////////////////////// // TargetStore @@ -2000,8 +1992,23 @@ type Datastore interface { // ///////////////////////////////////////////////////////////////////////////// // Android - SetAndroidEnabledAndConfigured(ctx context.Context, configured bool) error + AndroidDatastore +} + +type AndroidDatastore interface { + android.Datastore + AndroidHostLite(ctx context.Context, enterpriseSpecificID string) (*AndroidHost, error) + AppConfig(ctx context.Context) (*AppConfig, error) BulkSetAndroidHostsUnenrolled(ctx context.Context) error + DeleteMDMConfigAssetsByName(ctx context.Context, assetNames []MDMAssetName) error + GetAllMDMConfigAssetsByName(ctx context.Context, assetNames []MDMAssetName, + queryerContext sqlx.QueryerContext) (map[MDMAssetName]MDMConfigAsset, error) + InsertOrReplaceMDMConfigAsset(ctx context.Context, asset MDMConfigAsset) error + NewAndroidHost(ctx context.Context, host *AndroidHost) (*AndroidHost, error) + SetAndroidEnabledAndConfigured(ctx context.Context, configured bool) error + UpdateAndroidHost(ctx context.Context, host *AndroidHost, fromEnroll bool) error + UserOrDeletedUserByID(ctx context.Context, id uint) (*User, error) + VerifyEnrollSecret(ctx context.Context, secret string) (*EnrollSecret, error) } // MDMAppleStore wraps nanomdm's storage and adds methods to deal with diff --git a/server/mdm/android/service/enterprises_test.go b/server/mdm/android/service/enterprises_test.go index 2476103c8c..84a9d920e6 100644 --- a/server/mdm/android/service/enterprises_test.go +++ b/server/mdm/android/service/enterprises_test.go @@ -8,7 +8,6 @@ import ( "github.com/fleetdm/fleet/v4/server/authz" "github.com/fleetdm/fleet/v4/server/contexts/viewer" "github.com/fleetdm/fleet/v4/server/fleet" - "github.com/fleetdm/fleet/v4/server/mdm/android" android_mock "github.com/fleetdm/fleet/v4/server/mdm/android/mock" ds_mock "github.com/fleetdm/fleet/v4/server/mock" "github.com/fleetdm/fleet/v4/server/ptr" @@ -128,24 +127,20 @@ func checkAuthErr(t *testing.T, shouldFail bool, err error) { } } -func InitCommonDSMocks() *ds_mock.Store { - fleetDS := ds_mock.Store{} - ds := android_mock.Datastore{} - ds.InitCommonMocks() +func InitCommonDSMocks() fleet.AndroidDatastore { + ds := AndroidMockDS{} + ds.Datastore.InitCommonMocks() - fleetDS.GetAndroidDSFunc = func() android.Datastore { - return &ds - } - fleetDS.AppConfigFunc = func(_ context.Context) (*fleet.AppConfig, error) { + ds.Store.AppConfigFunc = func(_ context.Context) (*fleet.AppConfig, error) { return &fleet.AppConfig{}, nil } - fleetDS.SetAndroidEnabledAndConfiguredFunc = func(_ context.Context, configured bool) error { + ds.Store.SetAndroidEnabledAndConfiguredFunc = func(_ context.Context, configured bool) error { return nil } - fleetDS.UserOrDeletedUserByIDFunc = func(ctx context.Context, id uint) (*fleet.User, error) { + ds.Store.UserOrDeletedUserByIDFunc = func(ctx context.Context, id uint) (*fleet.User, error) { return &fleet.User{ID: id}, nil } - fleetDS.GetAllMDMConfigAssetsByNameFunc = func(ctx context.Context, assetNames []fleet.MDMAssetName, + ds.Store.GetAllMDMConfigAssetsByNameFunc = func(ctx context.Context, assetNames []fleet.MDMAssetName, queryerContext sqlx.QueryerContext) (map[fleet.MDMAssetName]fleet.MDMConfigAsset, error) { result := make(map[fleet.MDMAssetName]fleet.MDMConfigAsset, len(assetNames)) for _, name := range assetNames { @@ -153,16 +148,21 @@ func InitCommonDSMocks() *ds_mock.Store { } return result, nil } - fleetDS.InsertOrReplaceMDMConfigAssetFunc = func(ctx context.Context, asset fleet.MDMConfigAsset) error { + ds.Store.InsertOrReplaceMDMConfigAssetFunc = func(ctx context.Context, asset fleet.MDMConfigAsset) error { return nil } - fleetDS.DeleteMDMConfigAssetsByNameFunc = func(ctx context.Context, assetNames []fleet.MDMAssetName) error { + ds.Store.DeleteMDMConfigAssetsByNameFunc = func(ctx context.Context, assetNames []fleet.MDMAssetName) error { return nil } - fleetDS.BulkSetAndroidHostsUnenrolledFunc = func(ctx context.Context) error { + ds.Store.BulkSetAndroidHostsUnenrolledFunc = func(ctx context.Context) error { return nil } - return &fleetDS + return &ds +} + +type AndroidMockDS struct { + android_mock.Datastore + ds_mock.Store } type mockService struct { diff --git a/server/mdm/android/service/pubsub.go b/server/mdm/android/service/pubsub.go index 541fa8d57d..edf3ef7002 100644 --- a/server/mdm/android/service/pubsub.go +++ b/server/mdm/android/service/pubsub.go @@ -75,7 +75,7 @@ func (svc *Service) authenticatePubSub(ctx context.Context, token string) error // // Note: We could also check that the device belongs to our enterprise, for additional security. We would need an Android cached_mysql for that. // "name": "enterprises/LC044q09r2/devices/3dc9d72fbd517bbc", - assets, err := svc.fleetDS.GetAllMDMConfigAssetsByName(ctx, []fleet.MDMAssetName{fleet.MDMAssetAndroidPubSubToken}, nil) + assets, err := svc.ds.GetAllMDMConfigAssetsByName(ctx, []fleet.MDMAssetName{fleet.MDMAssetAndroidPubSubToken}, nil) switch { case fleet.IsNotFound(err): return fleet.NewAuthFailedError("missing Android PubSub token in Fleet") @@ -173,7 +173,7 @@ func (svc *Service) enrollHost(ctx context.Context, device *androidmanagement.De if host != nil { level.Debug(svc.logger).Log("msg", "The enrolling Android host is already present in Fleet. Updating team if needed", "device.name", device.Name, "device.enterpriseSpecificId", device.HardwareInfo.EnterpriseSpecificId) - enrollSecret, err := svc.fleetDS.VerifyEnrollSecret(ctx, device.EnrollmentTokenData) + enrollSecret, err := svc.ds.VerifyEnrollSecret(ctx, device.EnrollmentTokenData) if err != nil && !fleet.IsNotFound(err) { return ctxerr.Wrap(ctx, err, "verifying enroll secret") } @@ -251,7 +251,7 @@ func (svc *Service) updateHost(ctx context.Context, device *androidmanagement.De } host.SetNodeKey(device.HardwareInfo.EnterpriseSpecificId) - err = svc.fleetDS.UpdateAndroidHost(ctx, host, fromEnroll) + err = svc.ds.UpdateAndroidHost(ctx, host, fromEnroll) if err != nil { return ctxerr.Wrap(ctx, err, "enrolling Android host") } @@ -259,7 +259,7 @@ func (svc *Service) updateHost(ctx context.Context, device *androidmanagement.De } func (svc *Service) addNewHost(ctx context.Context, device *androidmanagement.Device) error { - enrollSecret, err := svc.fleetDS.VerifyEnrollSecret(ctx, device.EnrollmentTokenData) + enrollSecret, err := svc.ds.VerifyEnrollSecret(ctx, device.EnrollmentTokenData) if err != nil && !fleet.IsNotFound(err) { return ctxerr.Wrap(ctx, err, "verifying enroll secret") } @@ -301,7 +301,7 @@ func (svc *Service) addNewHost(ctx context.Context, device *androidmanagement.De host.Device.LastPolicySyncTime = ptr.Time(policySyncTime) } host.SetNodeKey(device.HardwareInfo.EnterpriseSpecificId) - _, err = svc.fleetDS.NewAndroidHost(ctx, host) + _, err = svc.ds.NewAndroidHost(ctx, host) if err != nil { return ctxerr.Wrap(ctx, err, "enrolling Android host") } @@ -314,7 +314,7 @@ func (svc *Service) getComputerName(device *androidmanagement.Device) string { } func (svc *Service) getHostIfPresent(ctx context.Context, enterpriseSpecificID string) (*fleet.AndroidHost, error) { - host, err := svc.fleetDS.AndroidHostLite(ctx, enterpriseSpecificID) + host, err := svc.ds.AndroidHostLite(ctx, enterpriseSpecificID) switch { case fleet.IsNotFound(err): return nil, nil diff --git a/server/mdm/android/service/service.go b/server/mdm/android/service/service.go index 041feb0cfd..f22faad381 100644 --- a/server/mdm/android/service/service.go +++ b/server/mdm/android/service/service.go @@ -30,8 +30,7 @@ const ( type Service struct { logger kitlog.Logger authz *authz.Authorizer - ds android.Datastore - fleetDS fleet.Datastore + ds fleet.AndroidDatastore proxy android.Proxy fleetSvc fleet.Service @@ -42,16 +41,16 @@ type Service struct { func NewService( ctx context.Context, logger kitlog.Logger, - fleetDS fleet.Datastore, + ds fleet.AndroidDatastore, fleetSvc fleet.Service, ) (android.Service, error) { prx := proxy.NewProxy(ctx, logger) - return NewServiceWithProxy(logger, fleetDS, prx, fleetSvc) + return NewServiceWithProxy(logger, ds, prx, fleetSvc) } func NewServiceWithProxy( logger kitlog.Logger, - fleetDS fleet.Datastore, + ds fleet.AndroidDatastore, proxy android.Proxy, fleetSvc fleet.Service, ) (android.Service, error) { @@ -63,8 +62,7 @@ func NewServiceWithProxy( return &Service{ logger: logger, authz: authorizer, - ds: fleetDS.GetAndroidDS(), - fleetDS: fleetDS, + ds: ds, proxy: proxy, fleetSvc: fleetSvc, SignupSSEInterval: DefaultSignupSSEInterval, @@ -129,7 +127,7 @@ func (svc *Service) EnterpriseSignup(ctx context.Context) (*android.SignupDetail } func (svc *Service) checkIfAndroidAlreadyConfigured(ctx context.Context) (*fleet.AppConfig, error) { - appConfig, err := svc.fleetDS.AppConfig(ctx) + appConfig, err := svc.ds.AppConfig(ctx) if err != nil { return nil, ctxerr.Wrap(ctx, err, "getting app config") } @@ -192,7 +190,7 @@ func (svc *Service) EnterpriseSignupCallback(ctx context.Context, signupToken st if err != nil { return ctxerr.Wrap(ctx, err, "generating pubsub token") } - err = svc.fleetDS.InsertOrReplaceMDMConfigAsset(ctx, fleet.MDMConfigAsset{ + err = svc.ds.InsertOrReplaceMDMConfigAsset(ctx, fleet.MDMConfigAsset{ Name: fleet.MDMAssetAndroidPubSubToken, Value: []byte(pubSubToken), }) @@ -261,12 +259,12 @@ func (svc *Service) EnterpriseSignupCallback(ctx context.Context, signupToken st return ctxerr.Wrap(ctx, err, "deleting temp enterprises") } - err = svc.fleetDS.SetAndroidEnabledAndConfigured(ctx, true) + err = svc.ds.SetAndroidEnabledAndConfigured(ctx, true) if err != nil { return ctxerr.Wrap(ctx, err, "setting android enabled and configured") } - user, err := svc.fleetDS.UserOrDeletedUserByID(ctx, enterprise.UserID) + user, err := svc.ds.UserOrDeletedUserByID(ctx, enterprise.UserID) switch { case fleet.IsNotFound(err): // This should never happen. @@ -341,12 +339,12 @@ func (svc *Service) DeleteEnterprise(ctx context.Context) error { return ctxerr.Wrap(ctx, err, "deleting enterprises") } - err = svc.fleetDS.SetAndroidEnabledAndConfigured(ctx, false) + err = svc.ds.SetAndroidEnabledAndConfigured(ctx, false) if err != nil { return ctxerr.Wrap(ctx, err, "clearing android enabled and configured") } - err = svc.fleetDS.BulkSetAndroidHostsUnenrolled(ctx) + err = svc.ds.BulkSetAndroidHostsUnenrolled(ctx) if err != nil { return ctxerr.Wrap(ctx, err, "bulk set android hosts as unenrolled") } @@ -355,7 +353,7 @@ func (svc *Service) DeleteEnterprise(ctx context.Context) error { return ctxerr.Wrap(ctx, err, "create activity for disabled Android MDM") } - err = svc.fleetDS.DeleteMDMConfigAssetsByName(ctx, []fleet.MDMAssetName{fleet.MDMAssetAndroidPubSubToken}) + err = svc.ds.DeleteMDMConfigAssetsByName(ctx, []fleet.MDMAssetName{fleet.MDMAssetAndroidPubSubToken}) if err != nil { return ctxerr.Wrap(ctx, err, "deleting pubsub token") } @@ -391,7 +389,7 @@ func (svc *Service) CreateEnrollmentToken(ctx context.Context, enrollSecret stri return nil, err } - _, err = svc.fleetDS.VerifyEnrollSecret(ctx, enrollSecret) + _, err = svc.ds.VerifyEnrollSecret(ctx, enrollSecret) switch { case fleet.IsNotFound(err): return nil, fleet.NewAuthFailedError("invalid secret") @@ -425,7 +423,7 @@ func (svc *Service) CreateEnrollmentToken(ctx context.Context, enrollSecret stri func (svc *Service) checkIfAndroidNotConfigured(ctx context.Context) (*fleet.AppConfig, error) { // This call uses cached_mysql implementation, so it's safe to call it multiple times - appConfig, err := svc.fleetDS.AppConfig(ctx) + appConfig, err := svc.ds.AppConfig(ctx) if err != nil { return nil, ctxerr.Wrap(ctx, err, "getting app config") } @@ -508,7 +506,7 @@ func (svc *Service) EnterpriseSignupSSE(ctx context.Context) (chan string, error } func (svc *Service) signupSSECheck(ctx context.Context, done chan string) bool { - appConfig, err := svc.fleetDS.AppConfig(ctx) + appConfig, err := svc.ds.AppConfig(ctx) if err != nil { done <- fmt.Sprintf("Error getting app config: %v", err) return true diff --git a/server/mdm/android/tests/enterprise/enterprise_test.go b/server/mdm/android/tests/enterprise/enterprise_test.go index 10e2fd0e5f..ba70aa2ff8 100644 --- a/server/mdm/android/tests/enterprise/enterprise_test.go +++ b/server/mdm/android/tests/enterprise/enterprise_test.go @@ -108,7 +108,7 @@ func (s *enterpriseTestSuite) TestEnterpriseSSE() { assert.Equal(s.T(), service.SignupSSESuccess, string(data)) // Test with error - s.WithServer.FleetDS.AppConfigFunc = func(_ context.Context) (*fleet.AppConfig, error) { + s.WithServer.DS.AppConfigFunc = func(_ context.Context) (*fleet.AppConfig, error) { return nil, assert.AnError } resp = s.Do("GET", "/api/v1/fleet/android_enterprise/signup_sse", nil, http.StatusOK) diff --git a/server/mdm/android/tests/testing_utils.go b/server/mdm/android/tests/testing_utils.go index 53cb8f25a8..91c0acbc0f 100644 --- a/server/mdm/android/tests/testing_utils.go +++ b/server/mdm/android/tests/testing_utils.go @@ -36,11 +36,15 @@ const ( EnterpriseID = "LC02k5wxw7" ) +type AndroidDSWithMock struct { + *mysql.Datastore + ds_mock.Store +} + type WithServer struct { suite.Suite Svc android.Service - DS *mysql.Datastore - FleetDS ds_mock.Store + DS AndroidDSWithMock FleetSvc mockService Server *httptest.Server Token string @@ -53,14 +57,14 @@ type WithServer struct { } func (ts *WithServer) SetupSuite(t *testing.T, dbName string) { - ts.DS = CreateNamedMySQLDS(t, dbName) + ts.DS.Datastore = CreateNamedMySQLDS(t, dbName) ts.CreateCommonDSMocks() ts.Proxy = proxy_mock.Proxy{} ts.createCommonProxyMocks(t) logger := kitlog.NewLogfmtLogger(os.Stdout) - svc, err := service.NewServiceWithProxy(logger, &ts.FleetDS, &ts.Proxy, &ts.FleetSvc) + svc, err := service.NewServiceWithProxy(logger, &ts.DS, &ts.Proxy, &ts.FleetSvc) require.NoError(t, err) ts.Svc = svc @@ -68,26 +72,23 @@ func (ts *WithServer) SetupSuite(t *testing.T, dbName string) { } func (ts *WithServer) CreateCommonDSMocks() { - ts.FleetDS.GetAndroidDSFunc = func() android.Datastore { - return ts.DS - } - ts.FleetDS.AppConfigFunc = func(_ context.Context) (*fleet.AppConfig, error) { + ts.DS.AppConfigFunc = func(_ context.Context) (*fleet.AppConfig, error) { // Create a copy to prevent race conditions ts.AppConfigMu.Lock() appConfigCopy := ts.AppConfig ts.AppConfigMu.Unlock() return &appConfigCopy, nil } - ts.FleetDS.SetAndroidEnabledAndConfiguredFunc = func(_ context.Context, configured bool) error { + ts.DS.SetAndroidEnabledAndConfiguredFunc = func(_ context.Context, configured bool) error { ts.AppConfigMu.Lock() ts.AppConfig.MDM.AndroidEnabledAndConfigured = configured ts.AppConfigMu.Unlock() return nil } - ts.FleetDS.UserOrDeletedUserByIDFunc = func(_ context.Context, id uint) (*fleet.User, error) { + ts.DS.UserOrDeletedUserByIDFunc = func(_ context.Context, id uint) (*fleet.User, error) { return &fleet.User{ID: id}, nil } - ts.FleetDS.GetAllMDMConfigAssetsByNameFunc = func(ctx context.Context, assetNames []fleet.MDMAssetName, + ts.DS.GetAllMDMConfigAssetsByNameFunc = func(ctx context.Context, assetNames []fleet.MDMAssetName, queryerContext sqlx.QueryerContext) (map[fleet.MDMAssetName]fleet.MDMConfigAsset, error) { result := make(map[fleet.MDMAssetName]fleet.MDMConfigAsset, len(assetNames)) for _, name := range assetNames { @@ -95,13 +96,13 @@ func (ts *WithServer) CreateCommonDSMocks() { } return result, nil } - ts.FleetDS.InsertOrReplaceMDMConfigAssetFunc = func(ctx context.Context, asset fleet.MDMConfigAsset) error { + ts.DS.InsertOrReplaceMDMConfigAssetFunc = func(ctx context.Context, asset fleet.MDMConfigAsset) error { return nil } - ts.FleetDS.DeleteMDMConfigAssetsByNameFunc = func(ctx context.Context, assetNames []fleet.MDMAssetName) error { + ts.DS.DeleteMDMConfigAssetsByNameFunc = func(ctx context.Context, assetNames []fleet.MDMAssetName) error { return nil } - ts.FleetDS.BulkSetAndroidHostsUnenrolledFunc = func(ctx context.Context) error { + ts.DS.BulkSetAndroidHostsUnenrolledFunc = func(ctx context.Context) error { return nil } } @@ -128,7 +129,7 @@ func (ts *WithServer) createCommonProxyMocks(t *testing.T) { } func (ts *WithServer) TearDownSuite() { - mysql.Close(ts.DS) + mysql.Close(ts.DS.Datastore) } type mockService struct { diff --git a/server/mock/datastore_mock.go b/server/mock/datastore_mock.go index 4b96f92174..9b02c8c84e 100644 --- a/server/mock/datastore_mock.go +++ b/server/mock/datastore_mock.go @@ -306,14 +306,6 @@ type OSVersionFunc func(ctx context.Context, osVersionID uint, teamFilter *fleet type UpdateOSVersionsFunc func(ctx context.Context) error -type GetAndroidDSFunc func() android.Datastore - -type NewAndroidHostFunc func(ctx context.Context, host *fleet.AndroidHost) (*fleet.AndroidHost, error) - -type UpdateAndroidHostFunc func(ctx context.Context, host *fleet.AndroidHost, fromEnroll bool) error - -type AndroidHostLiteFunc func(ctx context.Context, enterpriseSpecificID string) (*fleet.AndroidHost, error) - type CountHostsInTargetsFunc func(ctx context.Context, filter fleet.TeamFilter, targets fleet.HostTargets, now time.Time) (fleet.TargetMetrics, error) type HostIDsInTargetsFunc func(ctx context.Context, filter fleet.TeamFilter, targets fleet.HostTargets) ([]uint, error) @@ -1256,10 +1248,34 @@ type ExpandEmbeddedSecretsFunc func(ctx context.Context, document string) (strin type ExpandEmbeddedSecretsAndUpdatedAtFunc func(ctx context.Context, document string) (string, *time.Time, error) -type SetAndroidEnabledAndConfiguredFunc func(ctx context.Context, configured bool) error +type CreateEnterpriseFunc func(ctx context.Context, userID uint) (uint, error) + +type GetEnterpriseByIDFunc func(ctx context.Context, ID uint) (*android.EnterpriseDetails, error) + +type GetEnterpriseBySignupTokenFunc func(ctx context.Context, signupToken string) (*android.EnterpriseDetails, error) + +type GetEnterpriseFunc func(ctx context.Context) (*android.Enterprise, error) + +type UpdateEnterpriseFunc func(ctx context.Context, enterprise *android.EnterpriseDetails) error + +type DeleteAllEnterprisesFunc func(ctx context.Context) error + +type DeleteOtherEnterprisesFunc func(ctx context.Context, ID uint) error + +type CreateDeviceTxFunc func(ctx context.Context, tx sqlx.ExtContext, device *android.Device) (*android.Device, error) + +type UpdateDeviceTxFunc func(ctx context.Context, tx sqlx.ExtContext, device *android.Device) error + +type AndroidHostLiteFunc func(ctx context.Context, enterpriseSpecificID string) (*fleet.AndroidHost, error) type BulkSetAndroidHostsUnenrolledFunc func(ctx context.Context) error +type NewAndroidHostFunc func(ctx context.Context, host *fleet.AndroidHost) (*fleet.AndroidHost, error) + +type SetAndroidEnabledAndConfiguredFunc func(ctx context.Context, configured bool) error + +type UpdateAndroidHostFunc func(ctx context.Context, host *fleet.AndroidHost, fromEnroll bool) error + type DataStore struct { HealthCheckFunc HealthCheckFunc HealthCheckFuncInvoked bool @@ -1687,18 +1703,6 @@ type DataStore struct { UpdateOSVersionsFunc UpdateOSVersionsFunc UpdateOSVersionsFuncInvoked bool - GetAndroidDSFunc GetAndroidDSFunc - GetAndroidDSFuncInvoked bool - - NewAndroidHostFunc NewAndroidHostFunc - NewAndroidHostFuncInvoked bool - - UpdateAndroidHostFunc UpdateAndroidHostFunc - UpdateAndroidHostFuncInvoked bool - - AndroidHostLiteFunc AndroidHostLiteFunc - AndroidHostLiteFuncInvoked bool - CountHostsInTargetsFunc CountHostsInTargetsFunc CountHostsInTargetsFuncInvoked bool @@ -3112,12 +3116,48 @@ type DataStore struct { ExpandEmbeddedSecretsAndUpdatedAtFunc ExpandEmbeddedSecretsAndUpdatedAtFunc ExpandEmbeddedSecretsAndUpdatedAtFuncInvoked bool - SetAndroidEnabledAndConfiguredFunc SetAndroidEnabledAndConfiguredFunc - SetAndroidEnabledAndConfiguredFuncInvoked bool + CreateEnterpriseFunc CreateEnterpriseFunc + CreateEnterpriseFuncInvoked bool + + GetEnterpriseByIDFunc GetEnterpriseByIDFunc + GetEnterpriseByIDFuncInvoked bool + + GetEnterpriseBySignupTokenFunc GetEnterpriseBySignupTokenFunc + GetEnterpriseBySignupTokenFuncInvoked bool + + GetEnterpriseFunc GetEnterpriseFunc + GetEnterpriseFuncInvoked bool + + UpdateEnterpriseFunc UpdateEnterpriseFunc + UpdateEnterpriseFuncInvoked bool + + DeleteAllEnterprisesFunc DeleteAllEnterprisesFunc + DeleteAllEnterprisesFuncInvoked bool + + DeleteOtherEnterprisesFunc DeleteOtherEnterprisesFunc + DeleteOtherEnterprisesFuncInvoked bool + + CreateDeviceTxFunc CreateDeviceTxFunc + CreateDeviceTxFuncInvoked bool + + UpdateDeviceTxFunc UpdateDeviceTxFunc + UpdateDeviceTxFuncInvoked bool + + AndroidHostLiteFunc AndroidHostLiteFunc + AndroidHostLiteFuncInvoked bool BulkSetAndroidHostsUnenrolledFunc BulkSetAndroidHostsUnenrolledFunc BulkSetAndroidHostsUnenrolledFuncInvoked bool + NewAndroidHostFunc NewAndroidHostFunc + NewAndroidHostFuncInvoked bool + + SetAndroidEnabledAndConfiguredFunc SetAndroidEnabledAndConfiguredFunc + SetAndroidEnabledAndConfiguredFuncInvoked bool + + UpdateAndroidHostFunc UpdateAndroidHostFunc + UpdateAndroidHostFuncInvoked bool + mu sync.Mutex } @@ -4115,34 +4155,6 @@ func (s *DataStore) UpdateOSVersions(ctx context.Context) error { return s.UpdateOSVersionsFunc(ctx) } -func (s *DataStore) GetAndroidDS() android.Datastore { - s.mu.Lock() - s.GetAndroidDSFuncInvoked = true - s.mu.Unlock() - return s.GetAndroidDSFunc() -} - -func (s *DataStore) NewAndroidHost(ctx context.Context, host *fleet.AndroidHost) (*fleet.AndroidHost, error) { - s.mu.Lock() - s.NewAndroidHostFuncInvoked = true - s.mu.Unlock() - return s.NewAndroidHostFunc(ctx, host) -} - -func (s *DataStore) UpdateAndroidHost(ctx context.Context, host *fleet.AndroidHost, fromEnroll bool) error { - s.mu.Lock() - s.UpdateAndroidHostFuncInvoked = true - s.mu.Unlock() - return s.UpdateAndroidHostFunc(ctx, host, fromEnroll) -} - -func (s *DataStore) AndroidHostLite(ctx context.Context, enterpriseSpecificID string) (*fleet.AndroidHost, error) { - s.mu.Lock() - s.AndroidHostLiteFuncInvoked = true - s.mu.Unlock() - return s.AndroidHostLiteFunc(ctx, enterpriseSpecificID) -} - func (s *DataStore) CountHostsInTargets(ctx context.Context, filter fleet.TeamFilter, targets fleet.HostTargets, now time.Time) (fleet.TargetMetrics, error) { s.mu.Lock() s.CountHostsInTargetsFuncInvoked = true @@ -7440,11 +7452,74 @@ func (s *DataStore) ExpandEmbeddedSecretsAndUpdatedAt(ctx context.Context, docum return s.ExpandEmbeddedSecretsAndUpdatedAtFunc(ctx, document) } -func (s *DataStore) SetAndroidEnabledAndConfigured(ctx context.Context, configured bool) error { +func (s *DataStore) CreateEnterprise(ctx context.Context, userID uint) (uint, error) { s.mu.Lock() - s.SetAndroidEnabledAndConfiguredFuncInvoked = true + s.CreateEnterpriseFuncInvoked = true s.mu.Unlock() - return s.SetAndroidEnabledAndConfiguredFunc(ctx, configured) + return s.CreateEnterpriseFunc(ctx, userID) +} + +func (s *DataStore) GetEnterpriseByID(ctx context.Context, ID uint) (*android.EnterpriseDetails, error) { + s.mu.Lock() + s.GetEnterpriseByIDFuncInvoked = true + s.mu.Unlock() + return s.GetEnterpriseByIDFunc(ctx, ID) +} + +func (s *DataStore) GetEnterpriseBySignupToken(ctx context.Context, signupToken string) (*android.EnterpriseDetails, error) { + s.mu.Lock() + s.GetEnterpriseBySignupTokenFuncInvoked = true + s.mu.Unlock() + return s.GetEnterpriseBySignupTokenFunc(ctx, signupToken) +} + +func (s *DataStore) GetEnterprise(ctx context.Context) (*android.Enterprise, error) { + s.mu.Lock() + s.GetEnterpriseFuncInvoked = true + s.mu.Unlock() + return s.GetEnterpriseFunc(ctx) +} + +func (s *DataStore) UpdateEnterprise(ctx context.Context, enterprise *android.EnterpriseDetails) error { + s.mu.Lock() + s.UpdateEnterpriseFuncInvoked = true + s.mu.Unlock() + return s.UpdateEnterpriseFunc(ctx, enterprise) +} + +func (s *DataStore) DeleteAllEnterprises(ctx context.Context) error { + s.mu.Lock() + s.DeleteAllEnterprisesFuncInvoked = true + s.mu.Unlock() + return s.DeleteAllEnterprisesFunc(ctx) +} + +func (s *DataStore) DeleteOtherEnterprises(ctx context.Context, ID uint) error { + s.mu.Lock() + s.DeleteOtherEnterprisesFuncInvoked = true + s.mu.Unlock() + return s.DeleteOtherEnterprisesFunc(ctx, ID) +} + +func (s *DataStore) CreateDeviceTx(ctx context.Context, tx sqlx.ExtContext, device *android.Device) (*android.Device, error) { + s.mu.Lock() + s.CreateDeviceTxFuncInvoked = true + s.mu.Unlock() + return s.CreateDeviceTxFunc(ctx, tx, device) +} + +func (s *DataStore) UpdateDeviceTx(ctx context.Context, tx sqlx.ExtContext, device *android.Device) error { + s.mu.Lock() + s.UpdateDeviceTxFuncInvoked = true + s.mu.Unlock() + return s.UpdateDeviceTxFunc(ctx, tx, device) +} + +func (s *DataStore) AndroidHostLite(ctx context.Context, enterpriseSpecificID string) (*fleet.AndroidHost, error) { + s.mu.Lock() + s.AndroidHostLiteFuncInvoked = true + s.mu.Unlock() + return s.AndroidHostLiteFunc(ctx, enterpriseSpecificID) } func (s *DataStore) BulkSetAndroidHostsUnenrolled(ctx context.Context) error { @@ -7453,3 +7528,24 @@ func (s *DataStore) BulkSetAndroidHostsUnenrolled(ctx context.Context) error { s.mu.Unlock() return s.BulkSetAndroidHostsUnenrolledFunc(ctx) } + +func (s *DataStore) NewAndroidHost(ctx context.Context, host *fleet.AndroidHost) (*fleet.AndroidHost, error) { + s.mu.Lock() + s.NewAndroidHostFuncInvoked = true + s.mu.Unlock() + return s.NewAndroidHostFunc(ctx, host) +} + +func (s *DataStore) SetAndroidEnabledAndConfigured(ctx context.Context, configured bool) error { + s.mu.Lock() + s.SetAndroidEnabledAndConfiguredFuncInvoked = true + s.mu.Unlock() + return s.SetAndroidEnabledAndConfiguredFunc(ctx, configured) +} + +func (s *DataStore) UpdateAndroidHost(ctx context.Context, host *fleet.AndroidHost, fromEnroll bool) error { + s.mu.Lock() + s.UpdateAndroidHostFuncInvoked = true + s.mu.Unlock() + return s.UpdateAndroidHostFunc(ctx, host, fromEnroll) +} diff --git a/server/mock/mockimpl/impl.go b/server/mock/mockimpl/impl.go index d06425b1f4..fa389ef230 100644 --- a/server/mock/mockimpl/impl.go +++ b/server/mock/mockimpl/impl.go @@ -377,6 +377,17 @@ func main() { fatal(err) } + // Remove duplicates from fns + uniqueFns := make(map[string]Func, len(fns)) + dedupedFns := make([]Func, 0, len(fns)) + for _, fn := range fns { + if _, exists := uniqueFns[fn.Name]; !exists { + uniqueFns[fn.Name] = fn + dedupedFns = append(dedupedFns, fn) + } + } + fns = dedupedFns + src := genStubs(recv, fns) recName := strings.SplitN(recv, " ", 2) name := strings.TrimPrefix(recName[1], "*")