diff --git a/integration-test/collections/console_mps_apis.postman_collection.json b/integration-test/collections/console_mps_apis.postman_collection.json index 5943679a0..b21778421 100644 --- a/integration-test/collections/console_mps_apis.postman_collection.json +++ b/integration-test/collections/console_mps_apis.postman_collection.json @@ -1564,6 +1564,8 @@ " pm.expect(jsonData.totalCount).to.be.equal(0)\r", " pm.expect(jsonData.connectedCount).to.be.equal(0)\r", " pm.expect(jsonData.disconnectedCount).to.be.equal(0)\r", + " pm.expect(jsonData.activatedCount).to.be.equal(0)\r", + " pm.expect(jsonData.discoveredCount).to.be.equal(0)\r", " \r", "})" ], @@ -2268,6 +2270,110 @@ }, "response": [] }, + { + "name": "All Devices activated", + "event": [ + { + "listen": "test", + "script": { + "exec": [ + "pm.test(\"Status code is 200\", function () {\r", + " pm.response.to.have.status(200);\r", + "});" + ], + "type": "text/javascript", + "packages": {} + } + } + ], + "protocolProfileBehavior": { + "disableBodyPruning": true + }, + "request": { + "method": "GET", + "header": [], + "body": { + "mode": "raw", + "raw": "", + "options": { + "raw": { + "language": "json" + } + } + }, + "url": { + "raw": "{{protocol}}://{{host}}/api/v1/devices?activated=true", + "protocol": "{{protocol}}", + "host": [ + "{{host}}" + ], + "path": [ + "api", + "v1", + "devices" + ], + "query": [ + { + "key": "activated", + "value": "true" + } + ] + } + }, + "response": [] + }, + { + "name": "All Devices discovered", + "event": [ + { + "listen": "test", + "script": { + "exec": [ + "pm.test(\"Status code is 200\", function () {\r", + " pm.response.to.have.status(200);\r", + "});" + ], + "type": "text/javascript", + "packages": {} + } + } + ], + "protocolProfileBehavior": { + "disableBodyPruning": true + }, + "request": { + "method": "GET", + "header": [], + "body": { + "mode": "raw", + "raw": "", + "options": { + "raw": { + "language": "json" + } + } + }, + "url": { + "raw": "{{protocol}}://{{host}}/api/v1/devices?discovered=true", + "protocol": "{{protocol}}", + "host": [ + "{{host}}" + ], + "path": [ + "api", + "v1", + "devices" + ], + "query": [ + { + "key": "discovered", + "value": "true" + } + ] + } + }, + "response": [] + }, { "name": "All Devices with count set to true", "event": [ diff --git a/internal/controller/httpapi/v1/devices.go b/internal/controller/httpapi/v1/devices.go index 57485ec84..c7efec429 100644 --- a/internal/controller/httpapi/v1/devices.go +++ b/internal/controller/httpapi/v1/devices.go @@ -55,8 +55,18 @@ func (dr *deviceRoutes) getStats(c *gin.Context) { return } + activated, discovered, err := dr.t.GetDeviceStateCounts(c.Request.Context(), "") + if err != nil { + dr.l.Error(err, "http - devices - v1 - getStats") + ErrorResponse(c, err) + + return + } + countResponse := dto.DeviceStatResponse{ - TotalCount: count, + TotalCount: count, + ActivatedCount: activated, + DiscoveredCount: discovered, } c.JSON(http.StatusOK, countResponse) @@ -104,11 +114,15 @@ func (dr *deviceRoutes) get(c *gin.Context) { tags := c.Query("tags") hostname := c.Query("hostname") friendlyName := c.Query("friendlyName") + activated := c.Query("activated") + discovered := c.Query("discovered") var items []dto.Device var err error + ctx := c.Request.Context() + switch { case hostname != "": items, err = dr.getByColumnOrTags(c, "HostName", hostname, odata.Top, odata.Skip, "") @@ -119,8 +133,14 @@ func (dr *deviceRoutes) get(c *gin.Context) { case tags != "": items, err = dr.getByColumnOrTags(c, "Tags", tags, odata.Top, odata.Skip, "") + case activated == "true": + items, err = dr.t.GetActivated(ctx, odata.Top, odata.Skip, "") + + case discovered == "true": + items, err = dr.t.GetDiscovered(ctx, odata.Top, odata.Skip, "") + default: - items, err = dr.t.Get(c.Request.Context(), odata.Top, odata.Skip, "") + items, err = dr.t.Get(ctx, odata.Top, odata.Skip, "") } if err != nil { diff --git a/internal/controller/httpapi/v1/devices_test.go b/internal/controller/httpapi/v1/devices_test.go index a3319fc7e..fb64176e4 100644 --- a/internal/controller/httpapi/v1/devices_test.go +++ b/internal/controller/httpapi/v1/devices_test.go @@ -110,6 +110,40 @@ func TestDevicesRoutes(t *testing.T) { response: []dto.Device{{GUID: "guid", MPSUsername: "mpsusername", Username: "admin", Password: "password", ConnectionStatus: true, Hostname: "hostname"}}, expectedCode: http.StatusOK, }, + { + name: "get activated devices", + method: http.MethodGet, + url: "/api/v1/devices?activated=true", + mock: func(device *mocks.MockDeviceManagementFeature) { + device.EXPECT().GetActivated(context.Background(), 25, 0, "").Return([]dto.Device{{ + GUID: "guid", MPSUsername: "mpsusername", Username: "admin", Password: "password", ConnectionStatus: true, Hostname: "hostname", + }}, nil) + }, + response: []dto.Device{{GUID: "guid", MPSUsername: "mpsusername", Username: "admin", Password: "password", ConnectionStatus: true, Hostname: "hostname"}}, + expectedCode: http.StatusOK, + }, + { + name: "get discovered devices", + method: http.MethodGet, + url: "/api/v1/devices?discovered=true", + mock: func(device *mocks.MockDeviceManagementFeature) { + device.EXPECT().GetDiscovered(context.Background(), 25, 0, "").Return([]dto.Device{{ + GUID: "guid", MPSUsername: "mpsusername", Username: "admin", Password: "password", ConnectionStatus: true, Hostname: "hostname", + }}, nil) + }, + response: []dto.Device{{GUID: "guid", MPSUsername: "mpsusername", Username: "admin", Password: "password", ConnectionStatus: true, Hostname: "hostname"}}, + expectedCode: http.StatusOK, + }, + { + name: "get activated devices - failed", + method: http.MethodGet, + url: "/api/v1/devices?activated=true", + mock: func(device *mocks.MockDeviceManagementFeature) { + device.EXPECT().GetActivated(context.Background(), 25, 0, "").Return(nil, devices.ErrDatabase) + }, + response: devices.ErrDatabase, + expectedCode: http.StatusBadRequest, + }, { name: "get all devices - with count", method: http.MethodGet, @@ -317,10 +351,22 @@ func TestDevicesRoutes(t *testing.T) { url: "/api/v1/devices/stats", mock: func(device *mocks.MockDeviceManagementFeature) { device.EXPECT().GetCount(context.Background(), "").Return(5, nil) + device.EXPECT().GetDeviceStateCounts(context.Background(), "").Return(4, 1, nil) }, - response: dto.DeviceStatResponse{TotalCount: 5}, + response: dto.DeviceStatResponse{TotalCount: 5, ActivatedCount: 4, DiscoveredCount: 1}, expectedCode: http.StatusOK, }, + { + name: "get devices stats - failed", + method: http.MethodGet, + url: "/api/v1/devices/stats", + mock: func(device *mocks.MockDeviceManagementFeature) { + device.EXPECT().GetCount(context.Background(), "").Return(5, nil) + device.EXPECT().GetDeviceStateCounts(context.Background(), "").Return(0, 0, devices.ErrDatabase) + }, + response: devices.ErrDatabase, + expectedCode: http.StatusBadRequest, + }, } for _, tc := range tests { diff --git a/internal/controller/openapi/devices.go b/internal/controller/openapi/devices.go index f10d5f2b9..e4c08e080 100644 --- a/internal/controller/openapi/devices.go +++ b/internal/controller/openapi/devices.go @@ -60,6 +60,8 @@ func (f *FuegoAdapter) registerDeviceQueryRoutes() { fuego.OptionQuery("method", "Method to filter tags (any/all)"), fuego.OptionQuery("hostname", "Filter devices by host name"), fuego.OptionQuery("friendlyName", "Filter devices by friendly name"), + fuego.OptionQueryBool("activated", "Return devices activated into client or admin control mode"), + fuego.OptionQueryBool("discovered", "Return devices discovered on the network but not yet activated"), protectedRouteOptions(), ) @@ -205,6 +207,8 @@ func (f *FuegoAdapter) getDeviceStats(_ fuego.ContextNoBody) (dto.DeviceStatResp TotalCount: 5, ConnectedCount: 3, DisconnectedCount: 2, + ActivatedCount: 4, + DiscoveredCount: 1, }, nil } diff --git a/internal/entity/dto/v1/device.go b/internal/entity/dto/v1/device.go index c759649b1..e3ac28e2f 100644 --- a/internal/entity/dto/v1/device.go +++ b/internal/entity/dto/v1/device.go @@ -13,6 +13,8 @@ type DeviceStatResponse struct { TotalCount int `json:"totalCount"` ConnectedCount int `json:"connectedCount"` DisconnectedCount int `json:"disconnectedCount"` + ActivatedCount int `json:"activatedCount"` + DiscoveredCount int `json:"discoveredCount"` } type Device struct { ConnectionStatus bool `json:"connectionStatus"` diff --git a/internal/mocks/devicemanagement_mocks.go b/internal/mocks/devicemanagement_mocks.go index a5266d57d..93934bbe3 100644 --- a/internal/mocks/devicemanagement_mocks.go +++ b/internal/mocks/devicemanagement_mocks.go @@ -3,7 +3,7 @@ // // Generated by this command: // -// mockgen -source ./internal/usecase/devices/interfaces.go -package mocks -mock_names Repository=MockDeviceManagementRepository,Feature=MockDeviceManagementFeature +// mockgen -source ./internal/usecase/devices/interfaces.go -package mocks -mock_names Repository=MockDeviceManagementRepository,Feature=MockDeviceManagementFeature -destination ./internal/mocks/devicemanagement_mocks.go // // Package mocks is a generated GoMock package. @@ -307,6 +307,21 @@ func (mr *MockDeviceManagementRepositoryMockRecorder) Get(ctx, top, skip, tenant return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Get", reflect.TypeOf((*MockDeviceManagementRepository)(nil).Get), ctx, top, skip, tenantID) } +// GetActivated mocks base method. +func (m *MockDeviceManagementRepository) GetActivated(ctx context.Context, top, skip int, tenantID string) ([]entity.Device, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetActivated", ctx, top, skip, tenantID) + ret0, _ := ret[0].([]entity.Device) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetActivated indicates an expected call of GetActivated. +func (mr *MockDeviceManagementRepositoryMockRecorder) GetActivated(ctx, top, skip, tenantID any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetActivated", reflect.TypeOf((*MockDeviceManagementRepository)(nil).GetActivated), ctx, top, skip, tenantID) +} + // GetByColumn mocks base method. func (m *MockDeviceManagementRepository) GetByColumn(ctx context.Context, columnName, queryValue, tenantID string) ([]entity.Device, error) { m.ctrl.T.Helper() @@ -367,6 +382,37 @@ func (mr *MockDeviceManagementRepositoryMockRecorder) GetCount(arg0, arg1 any) * return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetCount", reflect.TypeOf((*MockDeviceManagementRepository)(nil).GetCount), arg0, arg1) } +// GetDeviceStateCounts mocks base method. +func (m *MockDeviceManagementRepository) GetDeviceStateCounts(ctx context.Context, tenantID string) (int, int, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetDeviceStateCounts", ctx, tenantID) + ret0, _ := ret[0].(int) + ret1, _ := ret[1].(int) + ret2, _ := ret[2].(error) + return ret0, ret1, ret2 +} + +// GetDeviceStateCounts indicates an expected call of GetDeviceStateCounts. +func (mr *MockDeviceManagementRepositoryMockRecorder) GetDeviceStateCounts(ctx, tenantID any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetDeviceStateCounts", reflect.TypeOf((*MockDeviceManagementRepository)(nil).GetDeviceStateCounts), ctx, tenantID) +} + +// GetDiscovered mocks base method. +func (m *MockDeviceManagementRepository) GetDiscovered(ctx context.Context, top, skip int, tenantID string) ([]entity.Device, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetDiscovered", ctx, top, skip, tenantID) + ret0, _ := ret[0].([]entity.Device) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetDiscovered indicates an expected call of GetDiscovered. +func (mr *MockDeviceManagementRepositoryMockRecorder) GetDiscovered(ctx, top, skip, tenantID any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetDiscovered", reflect.TypeOf((*MockDeviceManagementRepository)(nil).GetDiscovered), ctx, top, skip, tenantID) +} + // GetDistinctTags mocks base method. func (m *MockDeviceManagementRepository) GetDistinctTags(ctx context.Context, tenantID string) ([]string, error) { m.ctrl.T.Helper() @@ -594,6 +640,21 @@ func (mr *MockDeviceManagementFeatureMockRecorder) Get(ctx, top, skip, tenantID return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Get", reflect.TypeOf((*MockDeviceManagementFeature)(nil).Get), ctx, top, skip, tenantID) } +// GetActivated mocks base method. +func (m *MockDeviceManagementFeature) GetActivated(ctx context.Context, top, skip int, tenantID string) ([]dto.Device, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetActivated", ctx, top, skip, tenantID) + ret0, _ := ret[0].([]dto.Device) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetActivated indicates an expected call of GetActivated. +func (mr *MockDeviceManagementFeatureMockRecorder) GetActivated(ctx, top, skip, tenantID any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetActivated", reflect.TypeOf((*MockDeviceManagementFeature)(nil).GetActivated), ctx, top, skip, tenantID) +} + // GetAlarmOccurrences mocks base method. func (m *MockDeviceManagementFeature) GetAlarmOccurrences(ctx context.Context, guid string) ([]dto.AlarmClockOccurrence, error) { m.ctrl.T.Helper() @@ -729,6 +790,37 @@ func (mr *MockDeviceManagementFeatureMockRecorder) GetDeviceCertificate(c, guid return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetDeviceCertificate", reflect.TypeOf((*MockDeviceManagementFeature)(nil).GetDeviceCertificate), c, guid) } +// GetDeviceStateCounts mocks base method. +func (m *MockDeviceManagementFeature) GetDeviceStateCounts(ctx context.Context, tenantID string) (int, int, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetDeviceStateCounts", ctx, tenantID) + ret0, _ := ret[0].(int) + ret1, _ := ret[1].(int) + ret2, _ := ret[2].(error) + return ret0, ret1, ret2 +} + +// GetDeviceStateCounts indicates an expected call of GetDeviceStateCounts. +func (mr *MockDeviceManagementFeatureMockRecorder) GetDeviceStateCounts(ctx, tenantID any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetDeviceStateCounts", reflect.TypeOf((*MockDeviceManagementFeature)(nil).GetDeviceStateCounts), ctx, tenantID) +} + +// GetDiscovered mocks base method. +func (m *MockDeviceManagementFeature) GetDiscovered(ctx context.Context, top, skip int, tenantID string) ([]dto.Device, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetDiscovered", ctx, top, skip, tenantID) + ret0, _ := ret[0].([]dto.Device) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetDiscovered indicates an expected call of GetDiscovered. +func (mr *MockDeviceManagementFeatureMockRecorder) GetDiscovered(ctx, top, skip, tenantID any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetDiscovered", reflect.TypeOf((*MockDeviceManagementFeature)(nil).GetDiscovered), ctx, top, skip, tenantID) +} + // GetDiskInfo mocks base method. func (m *MockDeviceManagementFeature) GetDiskInfo(c context.Context, guid string) (dto.DiskInfo, error) { m.ctrl.T.Helper() diff --git a/internal/usecase/devices/interfaces.go b/internal/usecase/devices/interfaces.go index 931f49c43..71ac76192 100644 --- a/internal/usecase/devices/interfaces.go +++ b/internal/usecase/devices/interfaces.go @@ -46,6 +46,9 @@ type ( Update(ctx context.Context, d *entity.Device) (bool, error) Insert(ctx context.Context, d *entity.Device) (string, error) GetByColumn(ctx context.Context, columnName, queryValue, tenantID string) ([]entity.Device, error) + GetActivated(ctx context.Context, top, skip int, tenantID string) ([]entity.Device, error) + GetDiscovered(ctx context.Context, top, skip int, tenantID string) ([]entity.Device, error) + GetDeviceStateCounts(ctx context.Context, tenantID string) (activated, discovered int, err error) UpdateConnectionStatus(ctx context.Context, guid string, status bool) error UpdateLastSeen(ctx context.Context, guid string) error } @@ -62,6 +65,9 @@ type ( Update(ctx context.Context, d *dto.Device, fields map[string]bool) (*dto.Device, error) Insert(ctx context.Context, d *dto.Device) (*dto.Device, error) GetByColumn(ctx context.Context, columnName, queryValue, tenantID string) ([]dto.Device, error) + GetActivated(ctx context.Context, top, skip int, tenantID string) ([]dto.Device, error) + GetDiscovered(ctx context.Context, top, skip int, tenantID string) ([]dto.Device, error) + GetDeviceStateCounts(ctx context.Context, tenantID string) (activated, discovered int, err error) // Management Calls GetVersion(ctx context.Context, guid string) (dto.Version, dtov2.Version, error) GetFeatures(ctx context.Context, guid string) (dto.Features, dtov2.Features, error) diff --git a/internal/usecase/devices/repo.go b/internal/usecase/devices/repo.go index 28e067d52..804606534 100644 --- a/internal/usecase/devices/repo.go +++ b/internal/usecase/devices/repo.go @@ -76,6 +76,51 @@ func (uc *UseCase) GetByColumn(ctx context.Context, columnName, queryValue, tena return d1, nil } +// entitiesToDTOs converts a slice of device entities into their DTO representations. +func (uc *UseCase) entitiesToDTOs(data []entity.Device) ([]dto.Device, error) { + d1 := make([]dto.Device, len(data)) + + for i := range data { + tmpEntity := data[i] // create a new variable to avoid memory aliasing + + d, err := uc.entityToDTO(&tmpEntity) + if err != nil { + return nil, err + } + + d1[i] = *d + } + + return d1, nil +} + +func (uc *UseCase) GetActivated(ctx context.Context, top, skip int, tenantID string) ([]dto.Device, error) { + data, err := uc.repo.GetActivated(ctx, top, skip, tenantID) + if err != nil { + return nil, ErrDatabase.Wrap("GetActivated", "uc.repo.GetActivated", err) + } + + return uc.entitiesToDTOs(data) +} + +func (uc *UseCase) GetDiscovered(ctx context.Context, top, skip int, tenantID string) ([]dto.Device, error) { + data, err := uc.repo.GetDiscovered(ctx, top, skip, tenantID) + if err != nil { + return nil, ErrDatabase.Wrap("GetDiscovered", "uc.repo.GetDiscovered", err) + } + + return uc.entitiesToDTOs(data) +} + +func (uc *UseCase) GetDeviceStateCounts(ctx context.Context, tenantID string) (activated, discovered int, err error) { + activated, discovered, err = uc.repo.GetDeviceStateCounts(ctx, tenantID) + if err != nil { + return 0, 0, ErrDatabase.Wrap("GetDeviceStateCounts", "uc.repo.GetDeviceStateCounts", err) + } + + return activated, discovered, nil +} + func (uc *UseCase) GetByID(ctx context.Context, guid, tenantID string, includeSecrets bool) (*dto.Device, error) { data, err := uc.repo.GetByID(ctx, strings.ToLower(guid), tenantID) if err != nil { diff --git a/internal/usecase/devices/repo_test.go b/internal/usecase/devices/repo_test.go index c5ff390b5..4b6fd6b7f 100644 --- a/internal/usecase/devices/repo_test.go +++ b/internal/usecase/devices/repo_test.go @@ -184,6 +184,145 @@ func TestGet(t *testing.T) { } } +func TestGetActivated(t *testing.T) { + t.Parallel() + + testDevices := []entity.Device{{GUID: "guid-123", TenantID: "tenant-id-456"}} + testDeviceDTOs := []dto.Device{{GUID: "guid-123", TenantID: "tenant-id-456", Tags: nil}} + + tests := []testUsecase{ + { + name: "successful retrieval", + top: 10, + skip: 0, + tenantID: "tenant-id-456", + mock: func(repo *mocks.MockDeviceManagementRepository, _ *mocks.MockWSMAN) { + repo.EXPECT().GetActivated(context.Background(), 10, 0, "tenant-id-456").Return(testDevices, nil) + }, + res: testDeviceDTOs, + err: nil, + }, + { + name: "database error", + top: 5, + skip: 0, + tenantID: "tenant-id-456", + mock: func(repo *mocks.MockDeviceManagementRepository, _ *mocks.MockWSMAN) { + repo.EXPECT().GetActivated(context.Background(), 5, 0, "tenant-id-456").Return(nil, devices.ErrDatabase) + }, + res: []dto.Device(nil), + err: devices.ErrDatabase, + }, + } + + for _, tc := range tests { + tc := tc + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + useCase, repo, management := devicesTest(t) + + tc.mock(repo, management) + + results, err := useCase.GetActivated(context.Background(), tc.top, tc.skip, tc.tenantID) + + require.Equal(t, tc.res, results) + + if tc.err != nil { + require.Error(t, err) + require.Contains(t, err.Error(), tc.err.Error()) + } else { + require.NoError(t, err) + } + }) + } +} + +func TestGetDiscovered(t *testing.T) { + t.Parallel() + + testDevices := []entity.Device{{GUID: "guid-123", TenantID: "tenant-id-456"}} + testDeviceDTOs := []dto.Device{{GUID: "guid-123", TenantID: "tenant-id-456", Tags: nil}} + + tests := []testUsecase{ + { + name: "successful retrieval", + top: 10, + skip: 0, + tenantID: "tenant-id-456", + mock: func(repo *mocks.MockDeviceManagementRepository, _ *mocks.MockWSMAN) { + repo.EXPECT().GetDiscovered(context.Background(), 10, 0, "tenant-id-456").Return(testDevices, nil) + }, + res: testDeviceDTOs, + err: nil, + }, + { + name: "database error", + top: 5, + skip: 0, + tenantID: "tenant-id-456", + mock: func(repo *mocks.MockDeviceManagementRepository, _ *mocks.MockWSMAN) { + repo.EXPECT().GetDiscovered(context.Background(), 5, 0, "tenant-id-456").Return(nil, devices.ErrDatabase) + }, + res: []dto.Device(nil), + err: devices.ErrDatabase, + }, + } + + for _, tc := range tests { + tc := tc + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + useCase, repo, management := devicesTest(t) + + tc.mock(repo, management) + + results, err := useCase.GetDiscovered(context.Background(), tc.top, tc.skip, tc.tenantID) + + require.Equal(t, tc.res, results) + + if tc.err != nil { + require.Error(t, err) + require.Contains(t, err.Error(), tc.err.Error()) + } else { + require.NoError(t, err) + } + }) + } +} + +func TestGetDeviceStateCounts(t *testing.T) { + t.Parallel() + + t.Run("successful retrieval", func(t *testing.T) { + t.Parallel() + + useCase, repo, management := devicesTest(t) + _ = management + + repo.EXPECT().GetDeviceStateCounts(context.Background(), "tenant-id-456").Return(3, 1, nil) + + activated, discovered, err := useCase.GetDeviceStateCounts(context.Background(), "tenant-id-456") + require.NoError(t, err) + require.Equal(t, 3, activated) + require.Equal(t, 1, discovered) + }) + + t.Run("database error", func(t *testing.T) { + t.Parallel() + + useCase, repo, management := devicesTest(t) + _ = management + + repo.EXPECT().GetDeviceStateCounts(context.Background(), "tenant-id-456").Return(0, 0, devices.ErrDatabase) + + _, _, err := useCase.GetDeviceStateCounts(context.Background(), "tenant-id-456") + require.Error(t, err) + require.Contains(t, err.Error(), devices.ErrDatabase.Error()) + }) +} + func TestGetByID(t *testing.T) { t.Parallel() diff --git a/internal/usecase/nosqldb/mongo/device.go b/internal/usecase/nosqldb/mongo/device.go index 183ea7c6f..e3dde3b3b 100644 --- a/internal/usecase/nosqldb/mongo/device.go +++ b/internal/usecase/nosqldb/mongo/device.go @@ -345,3 +345,95 @@ func (r *DeviceRepo) GetByColumn(ctx context.Context, columnName, queryValue, te return devs, nil } + +// activatedFilter matches devices provisioned into a real AMT control mode +// (CCM/ACM): currentmode is present, non-empty, and not the "not activated" +// pre-provisioning string that rpc-go reports for un-provisioned devices. +func activatedFilter(tenantID string) bson.M { + return bson.M{ + fieldTenantID: tenantID, + fieldCurrentMode: bson.M{ + opExists: true, + opNin: bson.A{"", nil}, + opNot: bson.Regex{Pattern: "^not activated$", Options: "i"}, + }, + } +} + +// discoveredFilter is the exact complement of activatedFilter: a device is +// discovered (not yet activated) when currentmode is missing, null, empty, or +// the "not activated" pre-provisioning string. +func discoveredFilter(tenantID string) bson.M { + return bson.M{ + fieldTenantID: tenantID, + opOr: bson.A{ + bson.M{fieldCurrentMode: bson.M{opExists: false}}, + bson.M{fieldCurrentMode: bson.M{opIn: bson.A{"", nil}}}, + bson.M{fieldCurrentMode: bson.Regex{Pattern: "^not activated$", Options: "i"}}, + }, + } +} + +// GetActivated returns devices that have been provisioned into an AMT control mode. +func (r *DeviceRepo) GetActivated(ctx context.Context, top, skip int, tenantID string) ([]entity.Device, error) { + if tenantID != "" && !identifierRegex.MatchString(tenantID) { + return []entity.Device{}, nil + } + + return r.findFiltered(ctx, "GetActivated", activatedFilter(tenantID), top, skip) +} + +// GetDiscovered returns devices that have not yet been activated (currentmode empty/null/missing). +func (r *DeviceRepo) GetDiscovered(ctx context.Context, top, skip int, tenantID string) ([]entity.Device, error) { + if tenantID != "" && !identifierRegex.MatchString(tenantID) { + return []entity.Device{}, nil + } + + return r.findFiltered(ctx, "GetDiscovered", discoveredFilter(tenantID), top, skip) +} + +// GetDeviceStateCounts returns the number of activated and discovered devices for a tenant. +func (r *DeviceRepo) GetDeviceStateCounts(ctx context.Context, tenantID string) (activated, discovered int, err error) { + if tenantID != "" && !identifierRegex.MatchString(tenantID) { + return 0, 0, nil + } + + activatedCount, err := r.col.CountDocuments(ctx, activatedFilter(tenantID)) + if err != nil { + return 0, 0, errDeviceDatabase.Wrap("GetDeviceStateCounts", "CountDocuments", err) + } + + discoveredCount, err := r.col.CountDocuments(ctx, discoveredFilter(tenantID)) + if err != nil { + return 0, 0, errDeviceDatabase.Wrap("GetDeviceStateCounts", "CountDocuments", err) + } + + return int(activatedCount), int(discoveredCount), nil +} + +// findFiltered runs a paginated device query for an arbitrary filter (sorted by guid). +func (r *DeviceRepo) findFiltered(ctx context.Context, op string, filter bson.M, top, skip int) ([]entity.Device, error) { + limit := int64(DefaultTop) + if top > 0 { + limit = int64(top) + } + + offset := int64(0) + if skip > 0 { + offset = int64(skip) + } + + cur, err := r.col.Find(ctx, filter, + options.Find().SetSort(bson.D{{Key: fieldGUID, Value: 1}}).SetLimit(limit).SetSkip(offset)) + if err != nil { + return nil, errDeviceDatabase.Wrap(op, "Find", err) + } + defer cur.Close(ctx) + + devs := make([]entity.Device, 0) + if err := cur.All(ctx, &devs); err != nil { + return nil, errDeviceDatabase.Wrap(op, "Cursor.All", err) + } + + return devs, nil +} diff --git a/internal/usecase/nosqldb/mongo/device_test.go b/internal/usecase/nosqldb/mongo/device_test.go index 8fedd41ee..b2aa524e0 100644 --- a/internal/usecase/nosqldb/mongo/device_test.go +++ b/internal/usecase/nosqldb/mongo/device_test.go @@ -85,6 +85,75 @@ func TestDeviceRepo_Get(t *testing.T) { require.Len(t, rows, 2) } +func TestDeviceRepo_GetActivated(t *testing.T) { + t.Parallel() + + db, md := newMockedDB(t) + + md.AddResponses(findResponse( + "testdb."+mongo.CollectionDevices, + bson.D{{Key: "guid", Value: "g1"}, {Key: "currentmode", Value: "admin control mode"}, {Key: "tenantid", Value: "t1"}}, + )) + + repo := mongo.NewDeviceRepo(db) + + rows, err := repo.GetActivated(context.Background(), 10, 0, "t1") + require.NoError(t, err) + require.Len(t, rows, 1) + require.Equal(t, "g1", rows[0].GUID) +} + +func TestDeviceRepo_GetDiscovered(t *testing.T) { + t.Parallel() + + db, md := newMockedDB(t) + + md.AddResponses(findResponse( + "testdb."+mongo.CollectionDevices, + bson.D{{Key: "guid", Value: "g1"}, {Key: "currentmode", Value: ""}, {Key: "tenantid", Value: "t1"}}, + bson.D{{Key: "guid", Value: "g2"}, {Key: "currentmode", Value: "not activated"}, {Key: "tenantid", Value: "t1"}}, + )) + + repo := mongo.NewDeviceRepo(db) + + rows, err := repo.GetDiscovered(context.Background(), 10, 0, "t1") + require.NoError(t, err) + require.Len(t, rows, 2) + require.Equal(t, "g1", rows[0].GUID) + require.Equal(t, "g2", rows[1].GUID) +} + +func TestDeviceRepo_GetDeviceStateCounts(t *testing.T) { + t.Parallel() + + db, md := newMockedDB(t) + + // Two CountDocuments calls: activated first, then discovered. + md.AddResponses( + findResponse("testdb."+mongo.CollectionDevices, bson.D{{Key: "n", Value: int64(3)}}), + findResponse("testdb."+mongo.CollectionDevices, bson.D{{Key: "n", Value: int64(1)}}), + ) + + repo := mongo.NewDeviceRepo(db) + + activated, discovered, err := repo.GetDeviceStateCounts(context.Background(), "t1") + require.NoError(t, err) + require.Equal(t, 3, activated) + require.Equal(t, 1, discovered) +} + +func TestDeviceRepo_GetActivated_InvalidTenant(t *testing.T) { + t.Parallel() + + db, _ := newMockedDB(t) + + repo := mongo.NewDeviceRepo(db) + + rows, err := repo.GetActivated(context.Background(), 10, 0, "bad tenant!") + require.NoError(t, err) + require.Empty(t, rows) +} + // GetDistinctTags issues the `distinct` command, then de-duplicates and // trims tags in Go. The test asserts the post-driver in-process logic. func TestDeviceRepo_GetDistinctTags_DeduplicatesAcrossRows(t *testing.T) { diff --git a/internal/usecase/nosqldb/mongo/fields.go b/internal/usecase/nosqldb/mongo/fields.go index 65b55fd85..3ff8b9c69 100644 --- a/internal/usecase/nosqldb/mongo/fields.go +++ b/internal/usecase/nosqldb/mongo/fields.go @@ -14,9 +14,15 @@ const ( fieldWirelessProfileName = "wirelessprofilename" fieldPriority = "priority" fieldWiredInterface = "wiredinterface" + fieldCurrentMode = "currentmode" ) const ( - opSet = "$set" - opRegex = "$regex" + opSet = "$set" + opRegex = "$regex" + opNin = "$nin" + opExists = "$exists" + opIn = "$in" + opOr = "$or" + opNot = "$not" ) diff --git a/internal/usecase/sqldb/device.go b/internal/usecase/sqldb/device.go index d8fa3d4a2..30cf6e06e 100644 --- a/internal/usecase/sqldb/device.go +++ b/internal/usecase/sqldb/device.go @@ -25,6 +25,19 @@ var ( ErrDeviceNotUnique = repoerrors.NotUniqueError{Console: consoleerrors.CreateConsoleError("DeviceRepo")} ) +const ( + // activatedWhere matches devices provisioned into a real AMT control mode + // (CCM/ACM). rpc-go reports the interpreted control-mode string, so a + // pre-provisioning device sends the literal "not activated" — exclude it + // alongside NULL/empty legacy rows. + activatedWhere = "currentmode IS NOT NULL AND currentmode <> '' AND LOWER(currentmode) <> 'not activated'" + // discoveredWhere is the exact complement of activatedWhere: a device is + // discovered (not yet activated) when currentmode is NULL, empty, or the + // "not activated" pre-provisioning string. Parenthesised so the OR stays + // grouped when ANDed with the tenant clause. + discoveredWhere = "(currentmode IS NULL OR currentmode = '' OR LOWER(currentmode) = 'not activated')" +) + // New -. func NewDeviceRepo(database *db.SQL, log logger.Interface) *DeviceRepo { return &DeviceRepo{database, log} @@ -126,6 +139,126 @@ func (r *DeviceRepo) Get(_ context.Context, top, skip int, tenantID string) ([]e return devices, nil } +// GetActivated returns devices that have been provisioned into an AMT control mode. +func (r *DeviceRepo) GetActivated(_ context.Context, top, skip int, tenantID string) ([]entity.Device, error) { + return r.getFiltered("GetActivated", activatedWhere, nil, top, skip, tenantID) +} + +// GetDiscovered returns devices that have not yet been activated (currentmode empty/NULL). +func (r *DeviceRepo) GetDiscovered(_ context.Context, top, skip int, tenantID string) ([]entity.Device, error) { + return r.getFiltered("GetDiscovered", discoveredWhere, nil, top, skip, tenantID) +} + +// GetDeviceStateCounts returns the number of activated and discovered devices for a tenant. +func (r *DeviceRepo) GetDeviceStateCounts(_ context.Context, tenantID string) (activated, discovered int, err error) { + activated, err = r.countFiltered("GetDeviceStateCounts", activatedWhere, nil, tenantID) + if err != nil { + return 0, 0, err + } + + discovered, err = r.countFiltered("GetDeviceStateCounts", discoveredWhere, nil, tenantID) + if err != nil { + return 0, 0, err + } + + return activated, discovered, nil +} + +// getFiltered runs a paginated device query constrained by an extra WHERE clause. +func (r *DeviceRepo) getFiltered(op, whereClause string, whereArgs []any, top, skip int, tenantID string) ([]entity.Device, error) { + const defaultTop = 100 + + limitedTop := uint64(defaultTop) + if top > 0 { + limitedTop = uint64(top) + } + + limitedSkip := uint64(0) + if skip > 0 { + limitedSkip = uint64(skip) + } + + sqlQuery, args, err := r.Builder. + Select( + "guid", + "hostname", + "tags", + "mpsinstance", + "connectionstatus", + "mpsusername", + "tenantid", + "friendlyname", + "dnssuffix", + "deviceinfo", + "username", + "password", + "usetls", + "allowselfsigned", + "certhash", + ). + From("devices"). + Where("tenantid = ?", tenantID). + Where(whereClause, whereArgs...). + OrderBy("guid"). + Limit(limitedTop). + Offset(limitedSkip). + ToSql() + if err != nil { + return nil, ErrDeviceDatabase.Wrap(op, "r.Builder: ", err) + } + + rows, err := r.Pool.QueryContext(context.Background(), sqlQuery, args...) + if err != nil { + return nil, ErrDeviceDatabase.Wrap(op, "r.Pool.Query", err) + } + defer rows.Close() + + if rows.Err() != nil { + return nil, ErrDeviceDatabase.Wrap(op, "rows.Err", rows.Err()) + } + + devices := make([]entity.Device, 0) + + for rows.Next() { + d := entity.Device{} + + err = rows.Scan(&d.GUID, &d.Hostname, &d.Tags, &d.MPSInstance, &d.ConnectionStatus, &d.MPSUsername, &d.TenantID, &d.FriendlyName, &d.DNSSuffix, &d.DeviceInfo, &d.Username, &d.Password, &d.UseTLS, &d.AllowSelfSigned, &d.CertHash) + if err != nil { + return nil, ErrDeviceDatabase.Wrap(op, "rows.Scan: ", err) + } + + devices = append(devices, d) + } + + return devices, nil +} + +// countFiltered counts devices constrained by an extra WHERE clause. +func (r *DeviceRepo) countFiltered(op, whereClause string, whereArgs []any, tenantID string) (int, error) { + sqlQuery, args, err := r.Builder. + Select("COUNT(*)"). + From("devices"). + Where("tenantid = ?", tenantID). + Where(whereClause, whereArgs...). + ToSql() + if err != nil { + return 0, ErrDeviceDatabase.Wrap(op, "r.Builder: ", err) + } + + var count int + + err = r.Pool.QueryRowContext(context.Background(), sqlQuery, args...).Scan(&count) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return 0, nil + } + + return 0, ErrDeviceDatabase.Wrap(op, "r.Pool.QueryRow", err) + } + + return count, nil +} + // GetByID -. func (r *DeviceRepo) GetByID(_ context.Context, guid, tenantID string) (*entity.Device, error) { sqlQuery, _, err := r.Builder. diff --git a/internal/usecase/sqldb/device_test.go b/internal/usecase/sqldb/device_test.go index 214907133..b3585af8b 100644 --- a/internal/usecase/sqldb/device_test.go +++ b/internal/usecase/sqldb/device_test.go @@ -1294,3 +1294,117 @@ func TestDeviceRepo_UpdateLastSeen(t *testing.T) { }) } } + +// seedMode inserts a device row with the currentmode column set (nil -> SQL NULL). +func seedMode(t *testing.T, dbConn *sql.DB, guid, tenantID string, currentMode *string) { + t.Helper() + + _, err := dbConn.ExecContext(context.Background(), + `INSERT INTO devices (guid, tenantid, deviceinfo, currentmode) VALUES (?, ?, ?, ?)`, + guid, tenantID, "{}", currentMode) + require.NoError(t, err) +} + +func TestDeviceRepo_GetActivated(t *testing.T) { + t.Parallel() + + dbConn := setupDeviceTable(t) + defer dbConn.Close() + + // rpc-go reports interpreted control-mode strings (see utils.InterpretControlMode). + adminMode := "admin control mode" + clientMode := "client control mode" + notActivated := "not activated" + notActivatedCaps := "NOT ACTIVATED" + empty := "" + + seedMode(t, dbConn, "guid-admin", "tenant1", &adminMode) + seedMode(t, dbConn, "guid-client", "tenant1", &clientMode) + seedMode(t, dbConn, "guid-notactivated", "tenant1", ¬Activated) // pre-provisioning -> not activated + seedMode(t, dbConn, "guid-notactivated-caps", "tenant1", ¬ActivatedCaps) // case-insensitive exclude + seedMode(t, dbConn, "guid-empty", "tenant1", &empty) // empty -> not activated + seedMode(t, dbConn, "guid-nullmode", "tenant1", nil) // NULL -> not activated + seedMode(t, dbConn, "guid-othertenant", "tenant2", &adminMode) + + repo := sqldb.NewDeviceRepo(CreateSQLConfig(dbConn, false), mocks.NewMockLogger(nil)) + + result, err := repo.GetActivated(context.Background(), 0, 0, "tenant1") + require.NoError(t, err) + require.Len(t, result, 2) + + guids := []string{result[0].GUID, result[1].GUID} + assert.ElementsMatch(t, []string{"guid-admin", "guid-client"}, guids) +} + +func TestDeviceRepo_GetDiscovered(t *testing.T) { + t.Parallel() + + dbConn := setupDeviceTable(t) + defer dbConn.Close() + + adminMode := "admin control mode" + notActivated := "not activated" + notActivatedCaps := "NOT ACTIVATED" + empty := "" + seedMode(t, dbConn, "guid-empty", "tenant1", &empty) // empty -> discovered + seedMode(t, dbConn, "guid-null", "tenant1", nil) // NULL -> discovered + seedMode(t, dbConn, "guid-notactivated", "tenant1", ¬Activated) // pre-provisioning -> discovered + seedMode(t, dbConn, "guid-notactivated-caps", "tenant1", ¬ActivatedCaps) // case-insensitive + seedMode(t, dbConn, "guid-activated", "tenant1", &adminMode) + seedMode(t, dbConn, "guid-othertenant", "tenant2", ¬Activated) + + repo := sqldb.NewDeviceRepo(CreateSQLConfig(dbConn, false), mocks.NewMockLogger(nil)) + + result, err := repo.GetDiscovered(context.Background(), 0, 0, "tenant1") + require.NoError(t, err) + require.Len(t, result, 4) + + guids := make([]string, len(result)) + for i, d := range result { + guids[i] = d.GUID + } + + assert.ElementsMatch(t, []string{"guid-empty", "guid-null", "guid-notactivated", "guid-notactivated-caps"}, guids) +} + +func TestDeviceRepo_GetDeviceStateCounts(t *testing.T) { + t.Parallel() + + dbConn := setupDeviceTable(t) + defer dbConn.Close() + + adminMode := "admin control mode" + clientMode := "client control mode" + notActivated := "not activated" + empty := "" + + seedMode(t, dbConn, "guid-admin", "tenant1", &adminMode) + seedMode(t, dbConn, "guid-client", "tenant1", &clientMode) + seedMode(t, dbConn, "guid-notactivated", "tenant1", ¬Activated) + seedMode(t, dbConn, "guid-empty", "tenant1", &empty) + seedMode(t, dbConn, "guid-legacy", "tenant1", nil) + seedMode(t, dbConn, "guid-othertenant", "tenant2", &adminMode) + + repo := sqldb.NewDeviceRepo(CreateSQLConfig(dbConn, false), mocks.NewMockLogger(nil)) + + activated, discovered, err := repo.GetDeviceStateCounts(context.Background(), "tenant1") + require.NoError(t, err) + assert.Equal(t, 2, activated) + assert.Equal(t, 3, discovered) +} + +func TestDeviceRepo_GetDeviceStateCounts_Error(t *testing.T) { + t.Parallel() + + dbConn := setupDeviceTable(t) + defer dbConn.Close() + + repo := sqldb.NewDeviceRepo(CreateSQLConfig(dbConn, true), mocks.NewMockLogger(nil)) + + _, _, err := repo.GetDeviceStateCounts(context.Background(), "tenant1") + require.Error(t, err) + + var dbErr repoerrors.DatabaseError + + assert.True(t, errors.As(err, &dbErr)) +}